Merge remote-tracking branch 'origin/staging' into v2-architecture

This commit is contained in:
2026-03-22 14:33:59 -07:00
96 changed files with 12606 additions and 1383 deletions
+87 -3
View File
@@ -162,7 +162,7 @@ pub struct AgentDeps {
/// HTTP interceptor for trace recording/replay.
pub http_interceptor: Option<Arc<dyn crate::llm::recording::HttpInterceptor>>,
/// Audio transcription middleware for voice messages.
pub transcription: Option<Arc<crate::transcription::TranscriptionMiddleware>>,
pub transcription: Option<Arc<crate::llm::transcription::TranscriptionMiddleware>>,
/// Document text extraction middleware for PDF, DOCX, PPTX, etc.
pub document_extraction: Option<Arc<crate::document_extraction::DocumentExtractionMiddleware>>,
/// Sandbox readiness state for full-job routine dispatch.
@@ -1160,8 +1160,92 @@ impl Agent {
// Process based on submission type
let result = match submission {
Submission::UserInput { content } => {
self.process_user_input(message, session, thread_id, &content)
.await
let mut result = self
.process_user_input(message, session.clone(), thread_id, &content)
.await;
// Drain any messages queued during processing.
// Messages are merged (newline-separated) so the LLM receives
// full context from rapid consecutive inputs instead of
// processing each as a separate turn with partial context (#259).
//
// Only `Response` continues the drain — the user got a normal
// reply and there may be more queued messages to process.
//
// Everything else stops the loop:
// - `NeedApproval`: thread is blocked on user approval
// - `Interrupted`: turn was cancelled
// - `Ok`: control-command acknowledgment (including the "queued"
// ack returned when a message arrives during Processing)
// - `Error`: soft error — draining more messages after an error
// would produce confusing interleaved output
// - `Err(_)`: hard error
while let Ok(SubmissionResult::Response { content: outgoing }) = &result {
let merged = {
let mut sess = session.lock().await;
sess.threads
.get_mut(&thread_id)
.and_then(|t| t.drain_pending_messages())
};
let Some(next_content) = merged else {
break;
};
tracing::debug!(
thread_id = %thread_id,
merged_len = next_content.len(),
"Drain loop: processing merged queued messages"
);
// Send the completed turn's response before starting the next.
//
// Known limitations:
// - One-shot channels (HttpChannel) consume the response
// sender on the first respond() call keyed by msg.id.
// Subsequent calls (including the outer handler's final
// respond) are silently dropped. For one-shot channels
// only this intermediate response is delivered.
// - All drain-loop responses are routed via the original
// `message`, so channels that key routing on message
// identity will attribute every response to the first
// message. This is acceptable for the current
// single-user-per-thread model.
if let Err(e) = self
.channels
.respond(message, OutgoingResponse::text(outgoing.clone()))
.await
{
tracing::warn!(
thread_id = %thread_id,
"Failed to send intermediate drain-loop response: {e}"
);
}
// Process merged queued messages as a single turn.
// Use a message clone with cleared attachments so
// augment_with_attachments doesn't re-apply the original
// message's attachments to unrelated queued text.
let mut queued_msg = message.clone();
queued_msg.attachments.clear();
result = self
.process_user_input(&queued_msg, session.clone(), thread_id, &next_content)
.await;
// If processing failed, re-queue the drained content so it
// isn't lost. It will be picked up on the next successful turn.
if !matches!(&result, Ok(SubmissionResult::Response { .. })) {
let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&thread_id) {
thread.requeue_drained(next_content);
tracing::debug!(
thread_id = %thread_id,
"Re-queued drained content after non-Response result"
);
}
}
}
result
}
Submission::SystemCommand { command, args } => {
tracing::debug!(
+16 -3
View File
@@ -6,6 +6,7 @@
//! via the `LoopDelegate` trait.
use async_trait::async_trait;
use std::borrow::Cow;
use crate::agent::session::PendingApproval;
use crate::error::Error;
@@ -235,12 +236,12 @@ pub async fn run_agentic_loop(
///
/// `max` is a byte budget. The result is truncated at the last valid char
/// boundary at or before `max` bytes, so it is always valid UTF-8.
pub fn truncate_for_preview(s: &str, max: usize) -> String {
pub fn truncate_for_preview(s: &str, max: usize) -> Cow<'_, str> {
if s.len() <= max {
s.to_string()
Cow::Borrowed(s)
} else {
let end = crate::util::floor_char_boundary(s, max);
format!("{}...", &s[..end])
Cow::Owned(format!("{}...", &s[..end]))
}
}
@@ -597,12 +598,24 @@ mod tests {
assert_eq!(truncate_for_preview("hello", 10), "hello");
}
#[test]
fn test_truncate_short_string_borrows() {
let result = truncate_for_preview("hello", 10);
assert!(matches!(result, Cow::Borrowed("hello")));
}
#[test]
fn test_truncate_long_string_adds_ellipsis() {
let result = truncate_for_preview("hello world", 5);
assert_eq!(result, "hello...");
}
#[test]
fn test_truncate_long_string_owns() {
let result = truncate_for_preview("hello world", 5);
assert!(matches!(result, Cow::Owned(_)));
}
#[test]
fn test_truncate_multibyte_safe() {
let result = truncate_for_preview("café", 4);
+37 -20
View File
@@ -317,7 +317,7 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
.channels
.send_status(
&self.message.channel,
StatusUpdate::Thinking("Calling LLM...".into()),
StatusUpdate::Thinking(format!("Thinking (step {iteration})...")),
&self.message.metadata,
)
.await;
@@ -435,7 +435,7 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
.channels
.send_status(
&self.message.channel,
StatusUpdate::Thinking(format!("Executing {} tool(s)...", tool_calls.len())),
StatusUpdate::Thinking(contextual_tool_message(&tool_calls)),
&self.message.metadata,
)
.await;
@@ -845,11 +845,9 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
Ok(output) => {
let sanitized =
self.agent.safety().sanitize_tool_output(&tc.name, &output);
self.agent.safety().wrap_for_llm(
&tc.name,
&sanitized.content,
sanitized.was_modified,
)
self.agent
.safety()
.wrap_for_llm(&tc.name, &sanitized.content)
}
Err(e) => format!("Tool '{}' failed: {}", tc.name, e),
};
@@ -971,6 +969,30 @@ pub(super) fn check_auth_required(
Some((name, instructions))
}
/// Build a contextual thinking message based on tool names.
///
/// Instead of a generic "Executing 2 tool(s)..." this returns messages like
/// "Running command..." or "Fetching page..." for single-tool calls, falling
/// back to "Executing N tool(s)..." for multi-tool calls.
fn contextual_tool_message(tool_calls: &[crate::llm::ToolCall]) -> String {
if tool_calls.len() == 1 {
match tool_calls[0].name.as_str() {
"shell" => "Running command...".into(),
"web_fetch" => "Fetching page...".into(),
"memory_search" => "Searching memory...".into(),
"memory_write" => "Writing to memory...".into(),
"memory_read" => "Reading memory...".into(),
"http_request" => "Making HTTP request...".into(),
"file_read" => "Reading file...".into(),
"file_write" => "Writing file...".into(),
"json_transform" => "Transforming data...".into(),
name => format!("Running {name}..."),
}
} else {
format!("Executing {} tool(s)...", tool_calls.len())
}
}
/// Compact messages for retry after a context-length-exceeded error.
///
/// Keeps all `System` messages (which carry the system prompt and instructions),
@@ -1246,9 +1268,10 @@ mod tests {
#[test]
fn test_shell_destructive_command_requires_explicit_approval() {
// requires_explicit_approval() detects destructive commands that
// should return ApprovalRequirement::Always from ShellTool.
use crate::tools::builtin::shell::requires_explicit_approval;
// classify_command_risk() classifies destructive commands as High, which
// maps to ApprovalRequirement::Always in ShellTool::requires_approval().
use crate::tools::RiskLevel;
use crate::tools::builtin::shell::classify_command_risk;
let destructive_cmds = [
"rm -rf /tmp/test",
@@ -1256,20 +1279,14 @@ mod tests {
"git reset --hard HEAD~5",
];
for cmd in &destructive_cmds {
assert!(
requires_explicit_approval(cmd),
"'{}' should require explicit approval",
cmd
);
let r = classify_command_risk(cmd);
assert_eq!(r, RiskLevel::High, "'{}'", cmd); // safety: test code
}
let safe_cmds = ["git status", "cargo build", "ls -la"];
for cmd in &safe_cmds {
assert!(
!requires_explicit_approval(cmd),
"'{}' should not require explicit approval",
cmd
);
let r = classify_command_risk(cmd);
assert_ne!(r, RiskLevel::High, "'{}'", cmd); // safety: test code
}
}
+2 -2
View File
@@ -529,8 +529,8 @@ pub fn normalize_cron_expression(schedule: &str) -> String {
let trimmed = schedule.trim();
let fields: Vec<&str> = trimmed.split_whitespace().collect();
match fields.len() {
5 => format!("0 {} *", trimmed),
6 => format!("{} *", trimmed),
5 => format!("0 {} *", fields.join(" ")),
6 => format!("{} *", fields.join(" ")),
_ => trimmed.to_string(),
}
}
+2 -10
View File
@@ -1557,20 +1557,12 @@ async fn execute_lightweight_with_tools(
let result_content = match result {
Ok(output) => {
let sanitized = ctx.safety.sanitize_tool_output(&tc.name, &output);
ctx.safety.wrap_for_llm(
&tc.name,
&sanitized.content,
sanitized.was_modified,
)
ctx.safety.wrap_for_llm(&tc.name, &sanitized.content)
}
Err(e) => {
let error_msg = format!("Tool '{}' failed: {}", tc.name, e);
let sanitized = ctx.safety.sanitize_tool_output(&tc.name, &error_msg);
ctx.safety.wrap_for_llm(
&tc.name,
&sanitized.content,
sanitized.was_modified,
)
ctx.safety.wrap_for_llm(&tc.name, &sanitized.content)
}
};
+216 -2
View File
@@ -10,7 +10,7 @@
//! - Compaction: Summarize old turns to save context
//! - Resume: Continue from a saved checkpoint
use std::collections::{HashMap, HashSet};
use std::collections::{HashMap, HashSet, VecDeque};
use chrono::{DateTime, TimeDelta, Utc};
use serde::{Deserialize, Serialize};
@@ -222,8 +222,17 @@ pub struct Thread {
/// Pending auth token request (thread is in auth mode).
#[serde(default)]
pub pending_auth: Option<PendingAuth>,
/// Messages queued while the thread was processing a turn.
#[serde(default, skip_serializing_if = "VecDeque::is_empty")]
pub pending_messages: VecDeque<String>,
}
/// Maximum number of messages that can be queued while a thread is processing.
/// 10 merged messages can produce a large combined input for the LLM, but this
/// is acceptable for the personal assistant use case where a single user sends
/// rapid follow-ups. The drain loop processes them as one newline-delimited turn.
pub const MAX_PENDING_MESSAGES: usize = 10;
impl Thread {
/// Create a new thread.
pub fn new(session_id: Uuid) -> Self {
@@ -238,6 +247,7 @@ impl Thread {
metadata: serde_json::Value::Null,
pending_approval: None,
pending_auth: None,
pending_messages: VecDeque::new(),
}
}
@@ -254,6 +264,7 @@ impl Thread {
metadata: serde_json::Value::Null,
pending_approval: None,
pending_auth: None,
pending_messages: VecDeque::new(),
}
}
@@ -272,6 +283,47 @@ impl Thread {
self.turns.last_mut()
}
/// Queue a message for processing after the current turn completes.
/// Returns `false` if the queue is at capacity ([`MAX_PENDING_MESSAGES`]).
pub fn queue_message(&mut self, content: String) -> bool {
if self.pending_messages.len() >= MAX_PENDING_MESSAGES {
return false;
}
self.pending_messages.push_back(content);
self.updated_at = Utc::now();
true
}
/// Take the next pending message from the queue.
pub fn take_pending_message(&mut self) -> Option<String> {
self.pending_messages.pop_front()
}
/// Drain all pending messages from the queue.
/// Multiple messages are joined with newlines so the LLM receives
/// full context from rapid consecutive inputs (#259).
pub fn drain_pending_messages(&mut self) -> Option<String> {
if self.pending_messages.is_empty() {
return None;
}
let parts: Vec<String> = self.pending_messages.drain(..).collect();
self.updated_at = Utc::now();
Some(parts.join("\n"))
}
/// Re-queue previously drained content at the front of the queue.
/// Used to preserve user input when the drain loop fails to process
/// merged messages (soft error, hard error, interrupt).
///
/// This intentionally bypasses [`MAX_PENDING_MESSAGES`] — the content
/// was already counted against the cap before draining. The overshoot
/// is bounded to 1 entry (the re-queued merged string) plus any new
/// messages that arrived during the failed attempt.
pub fn requeue_drained(&mut self, content: String) {
self.pending_messages.push_front(content);
self.updated_at = Utc::now();
}
/// Start a new turn with user input.
pub fn start_turn(&mut self, user_input: impl Into<String>) -> &mut Turn {
let turn_number = self.turns.len();
@@ -335,11 +387,12 @@ impl Thread {
self.pending_auth.take()
}
/// Interrupt the current turn.
/// Interrupt the current turn and discard any queued messages.
pub fn interrupt(&mut self) {
if let Some(turn) = self.turns.last_mut() {
turn.interrupt();
}
self.pending_messages.clear();
self.state = ThreadState::Interrupted;
self.updated_at = Utc::now();
}
@@ -1392,4 +1445,165 @@ mod tests {
);
assert!(tool_result_content.ends_with("..."));
}
#[test]
fn test_thread_message_queue() {
let mut thread = Thread::new(Uuid::new_v4());
// Queue is initially empty
assert!(thread.pending_messages.is_empty());
assert!(thread.take_pending_message().is_none());
// Queue messages and verify FIFO ordering
assert!(thread.queue_message("first".to_string()));
assert!(thread.queue_message("second".to_string()));
assert!(thread.queue_message("third".to_string()));
assert_eq!(thread.pending_messages.len(), 3);
assert_eq!(thread.take_pending_message(), Some("first".to_string()));
assert_eq!(thread.take_pending_message(), Some("second".to_string()));
assert_eq!(thread.take_pending_message(), Some("third".to_string()));
assert!(thread.take_pending_message().is_none());
// Fill to capacity — all 10 should succeed
for i in 0..MAX_PENDING_MESSAGES {
assert!(thread.queue_message(format!("msg-{}", i)));
}
assert_eq!(thread.pending_messages.len(), MAX_PENDING_MESSAGES);
// 11th message rejected by queue_message itself
assert!(!thread.queue_message("overflow".to_string()));
assert_eq!(thread.pending_messages.len(), MAX_PENDING_MESSAGES);
// Drain and verify order
for i in 0..MAX_PENDING_MESSAGES {
assert_eq!(thread.take_pending_message(), Some(format!("msg-{}", i)));
}
assert!(thread.take_pending_message().is_none());
}
#[test]
fn test_thread_message_queue_serialization() {
let mut thread = Thread::new(Uuid::new_v4());
// Empty queue should not appear in serialization (skip_serializing_if)
let json = serde_json::to_string(&thread).unwrap();
assert!(!json.contains("pending_messages"));
// Non-empty queue should serialize and deserialize
thread.queue_message("queued msg".to_string());
let json = serde_json::to_string(&thread).unwrap();
assert!(json.contains("pending_messages"));
assert!(json.contains("queued msg"));
let restored: Thread = serde_json::from_str(&json).unwrap();
assert_eq!(restored.pending_messages.len(), 1);
assert_eq!(restored.pending_messages[0], "queued msg");
}
#[test]
fn test_thread_message_queue_default_on_old_data() {
// Deserialization of old data without pending_messages should default to empty
let thread = Thread::new(Uuid::new_v4());
let json = serde_json::to_string(&thread).unwrap();
// The field is absent (skip_serializing_if), simulating old data
assert!(!json.contains("pending_messages"));
let restored: Thread = serde_json::from_str(&json).unwrap();
assert!(restored.pending_messages.is_empty());
}
#[test]
fn test_interrupt_clears_pending_messages() {
let mut thread = Thread::new(Uuid::new_v4());
// Start a turn so there's something to interrupt
thread.start_turn("initial input");
// Queue several messages while "processing"
thread.queue_message("queued-1".to_string());
thread.queue_message("queued-2".to_string());
thread.queue_message("queued-3".to_string());
assert_eq!(thread.pending_messages.len(), 3);
// Interrupt should clear the queue
thread.interrupt();
assert!(thread.pending_messages.is_empty());
assert_eq!(thread.state, ThreadState::Interrupted);
}
#[test]
fn test_thread_state_idle_after_full_drain() {
let mut thread = Thread::new(Uuid::new_v4());
// Simulate a full drain cycle: start turn, queue messages, complete turn,
// then drain all queued messages as a single merged turn (#259).
thread.start_turn("turn 1");
assert_eq!(thread.state, ThreadState::Processing);
thread.queue_message("queued-a".to_string());
thread.queue_message("queued-b".to_string());
// Complete the turn (simulates process_user_input finishing)
thread.complete_turn("response 1");
assert_eq!(thread.state, ThreadState::Idle);
// Drain: merge all queued messages and process as a single turn
let merged = thread.drain_pending_messages().unwrap();
assert_eq!(merged, "queued-a\nqueued-b");
thread.start_turn(&merged);
thread.complete_turn("response for merged");
// Queue is fully drained, thread is idle
assert!(thread.drain_pending_messages().is_none());
assert!(thread.pending_messages.is_empty());
assert_eq!(thread.state, ThreadState::Idle);
}
#[test]
fn test_drain_pending_messages_merges_with_newlines() {
let mut thread = Thread::new(Uuid::new_v4());
// Empty queue returns None
assert!(thread.drain_pending_messages().is_none());
// Single message returned as-is (no trailing newline)
thread.queue_message("only one".to_string());
assert_eq!(
thread.drain_pending_messages(),
Some("only one".to_string()),
);
assert!(thread.pending_messages.is_empty());
// Multiple messages joined with newlines
thread.queue_message("hey".to_string());
thread.queue_message("can you check the server".to_string());
thread.queue_message("it started 10 min ago".to_string());
assert_eq!(
thread.drain_pending_messages(),
Some("hey\ncan you check the server\nit started 10 min ago".to_string()),
);
assert!(thread.pending_messages.is_empty());
// Queue is empty after drain
assert!(thread.drain_pending_messages().is_none());
}
#[test]
fn test_requeue_drained_preserves_content_at_front() {
let mut thread = Thread::new(Uuid::new_v4());
// Re-queue into empty queue
thread.requeue_drained("failed batch".to_string());
assert_eq!(thread.pending_messages.len(), 1);
assert_eq!(thread.pending_messages[0], "failed batch");
// New messages go behind the re-queued content
thread.queue_message("new msg".to_string());
assert_eq!(thread.pending_messages.len(), 2);
// Drain should return re-queued content first (front of queue)
let merged = thread.drain_pending_messages().unwrap();
assert_eq!(merged, "failed batch\nnew msg");
}
}
+201 -9
View File
@@ -14,7 +14,7 @@ use crate::agent::compaction::ContextCompactor;
use crate::agent::dispatcher::{
AgenticLoopResult, check_auth_required, execute_chat_tool_standalone, parse_auth_result,
};
use crate::agent::session::{PendingApproval, Session, ThreadState};
use crate::agent::session::{MAX_PENDING_MESSAGES, PendingApproval, Session, ThreadState};
use crate::agent::submission::SubmissionResult;
use crate::channels::web::util::truncate_preview;
use crate::channels::{IncomingMessage, StatusUpdate};
@@ -211,14 +211,72 @@ impl Agent {
// Check thread state
match thread_state {
ThreadState::Processing => {
tracing::warn!(
message_id = %message.id,
thread_id = %thread_id,
"Thread is processing, rejecting new input"
);
return Ok(SubmissionResult::error(
"Turn in progress. Use /interrupt to cancel.",
));
let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&thread_id) {
// Re-check state under lock — the turn may have completed
// between the snapshot read and this mutable lock acquisition.
if thread.state == ThreadState::Processing {
// Reject messages with attachments — the queue stores
// text only, so attachments would be silently dropped.
if !message.attachments.is_empty() {
return Ok(SubmissionResult::error(
"Cannot queue messages with attachments while a turn is processing. \
Please resend after the current turn completes.",
));
}
// Run the same safety checks that the normal path applies
// (validation, policy, secret scan) so that blocked content
// is never stored in pending_messages or serialized.
let validation = self.safety().validate_input(content);
if !validation.is_valid {
let details = validation
.errors
.iter()
.map(|e| format!("{}: {}", e.field, e.message))
.collect::<Vec<_>>()
.join("; ");
return Ok(SubmissionResult::error(format!(
"Input rejected by safety validation: {details}",
)));
}
let violations = self.safety().check_policy(content);
if violations
.iter()
.any(|rule| rule.action == crate::safety::PolicyAction::Block)
{
return Ok(SubmissionResult::error("Input rejected by safety policy."));
}
if let Some(warning) = self.safety().scan_inbound_for_secrets(content) {
tracing::warn!(
user = %message.user_id,
channel = %message.channel,
"Queued message blocked: contains leaked secret"
);
return Ok(SubmissionResult::error(warning));
}
if !thread.queue_message(content.to_string()) {
return Ok(SubmissionResult::error(format!(
"Message queue full ({MAX_PENDING_MESSAGES}). Wait for the current turn to complete.",
)));
}
// Return `Ok` (not `Response`) so the drain loop in
// agent_loop.rs breaks — `Ok` signals a control
// acknowledgment, not a completed LLM turn.
return Ok(SubmissionResult::Ok {
message: Some(
"Message queued — will be processed after the current turn.".into(),
),
});
}
// State changed (turn completed) — fall through to process normally.
// NOTE: `sess` (the Mutex guard) is dropped at the end of
// this `Processing` match arm, releasing the session lock
// before the rest of process_user_input runs. No deadlock.
} else {
return Ok(SubmissionResult::error("Thread no longer exists."));
}
}
ThreadState::AwaitingApproval => {
tracing::warn!(
@@ -498,6 +556,33 @@ impl Agent {
.await;
}
// Emit per-turn cost summary
{
let usage = self.cost_guard().model_usage().await;
let (total_in, total_out, total_cost) =
usage
.values()
.fold((0u64, 0u64, rust_decimal::Decimal::ZERO), |acc, m| {
(
acc.0 + m.input_tokens,
acc.1 + m.output_tokens,
acc.2 + m.cost,
)
});
let _ = self
.channels
.send_status(
&message.channel,
StatusUpdate::TurnCost {
input_tokens: total_in,
output_tokens: total_out,
cost_usd: format!("${:.4}", total_cost),
},
&message.metadata,
)
.await;
}
Ok(SubmissionResult::response(response))
}
Ok(AgenticLoopResult::NeedApproval { pending }) => {
@@ -849,6 +934,7 @@ impl Agent {
.get_mut(&thread_id)
.ok_or_else(|| Error::from(crate::error::JobError::NotFound { id: thread_id }))?;
thread.turns.clear();
thread.pending_messages.clear();
thread.state = ThreadState::Idle;
// Clear undo history too
@@ -2012,6 +2098,112 @@ mod tests {
}
}
#[test]
fn test_queue_cap_rejects_at_capacity() {
use crate::agent::session::{MAX_PENDING_MESSAGES, Thread, ThreadState};
use uuid::Uuid;
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("processing something");
assert_eq!(thread.state, ThreadState::Processing);
// Fill the queue to the cap
for i in 0..MAX_PENDING_MESSAGES {
assert!(thread.queue_message(format!("msg-{}", i)));
}
assert_eq!(thread.pending_messages.len(), MAX_PENDING_MESSAGES);
// The next message should be rejected by queue_message
assert!(!thread.queue_message("overflow".to_string()));
assert_eq!(thread.pending_messages.len(), MAX_PENDING_MESSAGES);
// Verify all drain in FIFO order
for i in 0..MAX_PENDING_MESSAGES {
assert_eq!(thread.take_pending_message(), Some(format!("msg-{}", i)));
}
assert!(thread.take_pending_message().is_none());
}
#[test]
fn test_clear_clears_pending_messages() {
use crate::agent::session::{Thread, ThreadState};
use uuid::Uuid;
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("processing");
thread.queue_message("pending-1".to_string());
thread.queue_message("pending-2".to_string());
assert_eq!(thread.pending_messages.len(), 2);
// Simulate what process_clear does: clear turns and pending_messages
thread.turns.clear();
thread.pending_messages.clear();
thread.state = ThreadState::Idle;
assert!(thread.pending_messages.is_empty());
assert!(thread.turns.is_empty());
assert_eq!(thread.state, ThreadState::Idle);
}
#[test]
fn test_processing_arm_thread_gone_returns_error() {
// Regression: if the thread disappears between the state snapshot and the
// mutable lock, the Processing arm must return an error — not a false
// "queued" acknowledgment.
//
// Exercises the exact branch at the `else` of
// `if let Some(thread) = sess.threads.get_mut(&thread_id)`.
use crate::agent::session::{Session, Thread, ThreadState};
use uuid::Uuid;
let thread_id = Uuid::new_v4();
let session_id = Uuid::new_v4();
let mut thread = Thread::with_id(thread_id, session_id);
thread.start_turn("working");
assert_eq!(thread.state, ThreadState::Processing);
let mut session = Session::new("test-user");
session.threads.insert(thread_id, thread);
// Simulate the thread disappearing (e.g., /clear racing with queue)
session.threads.remove(&thread_id);
// The Processing arm re-locks and calls get_mut — must get None.
assert!(session.threads.get_mut(&thread_id).is_none());
// Nothing was queued anywhere — the removed thread's queue is gone.
}
#[test]
fn test_processing_arm_state_changed_does_not_queue() {
// Regression: if the thread transitions from Processing to Idle between
// the state snapshot and the mutable lock, the message must NOT be queued.
// Instead the Processing arm falls through to normal processing.
//
// Exercises the `if thread.state == ThreadState::Processing` re-check.
use crate::agent::session::{Session, Thread, ThreadState};
use uuid::Uuid;
let thread_id = Uuid::new_v4();
let session_id = Uuid::new_v4();
let mut thread = Thread::with_id(thread_id, session_id);
thread.start_turn("working");
assert_eq!(thread.state, ThreadState::Processing);
// Simulate the turn completing between snapshot and re-lock
thread.complete_turn("done");
assert_eq!(thread.state, ThreadState::Idle);
let mut session = Session::new("test-user");
session.threads.insert(thread_id, thread);
// Re-check under lock: state is Idle, so queue_message must NOT be called.
let t = session.threads.get_mut(&thread_id).unwrap();
assert_ne!(t.state, ThreadState::Processing);
// Verify nothing was queued — the fall-through path doesn't touch the queue.
assert!(t.pending_messages.is_empty());
}
// Helper function to extract the approval message without needing a full Agent instance
fn extract_approval_message(
session: &crate::agent::session::Session,
+42 -12
View File
@@ -386,7 +386,7 @@ impl AppBuilder {
let b = tools
.register_builder_tool(llm.clone(), Some(self.config.builder.to_builder_config()))
.await;
tracing::info!("Builder mode enabled");
tracing::debug!("Builder mode enabled");
Some(b)
} else {
None
@@ -536,7 +536,7 @@ impl AppBuilder {
server_name,
e
);
return;
return None;
}
};
@@ -553,6 +553,10 @@ impl AppBuilder {
tool_count,
server_name
);
return Some((
server_name,
Arc::new(client),
));
}
Err(e) => {
tracing::warn!(
@@ -583,14 +587,27 @@ impl AppBuilder {
}
}
}
None
});
}
let mut startup_clients = Vec::new();
while let Some(result) = join_set.join_next().await {
if let Err(e) = result {
tracing::warn!("MCP server loading task panicked: {}", e);
match result {
Ok(Some(client_pair)) => {
startup_clients.push(client_pair);
}
Ok(None) => {}
Err(e) => {
if e.is_panic() {
tracing::error!("MCP server loading task panicked: {}", e);
} else {
tracing::warn!("MCP server loading task failed: {}", e);
}
}
}
}
return startup_clients;
}
Err(e) => {
if matches!(
@@ -608,10 +625,12 @@ impl AppBuilder {
}
}
}
Vec::new()
}
};
let (dev_loaded_tool_names, _) = tokio::join!(wasm_tools_future, mcp_servers_future);
let (dev_loaded_tool_names, startup_mcp_clients) =
tokio::join!(wasm_tools_future, mcp_servers_future);
// Load registry catalog entries for extension discovery
let mut catalog_entries = match crate::registry::RegistryCatalog::load_or_embedded() {
@@ -673,6 +692,17 @@ impl AppBuilder {
));
tools.register_extension_tools(Arc::clone(&manager));
tracing::debug!("Extension manager initialized with in-chat discovery tools");
if !startup_mcp_clients.is_empty() {
tracing::info!(
count = startup_mcp_clients.len(),
"Injecting startup MCP clients into extension manager"
);
for (name, client) in startup_mcp_clients {
manager.inject_mcp_client(name, client).await;
}
}
Some(manager)
};
@@ -699,13 +729,13 @@ impl AppBuilder {
self.init_database().await?;
self.init_secrets().await?;
// Post-init validation: if a non-nearai backend was selected but
// credentials were never resolved (deferred resolution found no keys),
// fail early with a clear error instead of a confusing runtime failure.
if self.config.llm.backend != "nearai"
&& self.config.llm.backend != "bedrock"
&& self.config.llm.backend != "openai_codex"
&& self.config.llm.provider.is_none()
// Post-init validation: backends with dedicated config (nearai, gemini_oauth,
// bedrock, openai_codex) handle their own credential resolution. For registry-based
// backends, fail early if no provider config was resolved.
if !matches!(
self.config.llm.backend.as_str(),
"nearai" | "gemini_oauth" | "bedrock" | "openai_codex"
) && self.config.llm.provider.is_none()
{
let backend = &self.config.llm.backend;
anyhow::bail!(
+188 -93
View File
@@ -1,8 +1,11 @@
//! Boot screen displayed after all initialization completes.
//!
//! Shows a polished ANSI-styled status panel summarizing the agent's runtime
//! state: model, database, tool count, enabled features, active channels,
//! and the gateway URL.
//! Shows a compact ANSI-styled status panel with three tiers:
//! - **Tier 1 (always):** Name + version, model + backend.
//! - **Tier 2 (conditional):** Gateway URL, tunnel URL, non-default channels.
//! - **Tier 3 (removed):** Database, tool count, features → use `ironclaw status`.
use crate::cli::fmt;
/// All displayable fields for the boot screen.
pub struct BootInfo {
@@ -29,112 +32,76 @@ pub struct BootInfo {
pub tunnel_url: Option<String>,
/// Provider name for the managed tunnel (e.g., "ngrok").
pub tunnel_provider: Option<String>,
/// Time elapsed during startup. Shown at the bottom when present.
pub startup_elapsed: Option<std::time::Duration>,
}
/// Print the boot screen to stdout.
pub fn print_boot_screen(info: &BootInfo) {
// ANSI codes matching existing REPL palette
let bold = "\x1b[1m";
let cyan = "\x1b[36m";
let dim = "\x1b[90m";
let yellow = "\x1b[33m";
let yellow_underline = "\x1b[33;4m";
let reset = "\x1b[0m";
const KW: usize = 10;
let border = format!(" {dim}{}{reset}", "\u{2576}".repeat(58));
/// Print the boot screen to stdout.
///
/// **Tier 1 (always):** Name + version, model + backend.
/// **Tier 2 (conditional):** Gateway URL, tunnel URL, non-default channels.
/// **Tier 3 (removed):** Database, tool count, features — use `ironclaw status`.
pub fn print_boot_screen(info: &BootInfo) {
let border = format!(" {}", fmt::separator(58));
println!();
println!("{border}");
println!();
println!(" {bold}{}{reset} v{}", info.agent_name, info.version);
// ── Tier 1: always shown ──────────────────────────────────────────
println!(
" {}{}{} v{}",
fmt::bold(),
info.agent_name,
fmt::reset(),
info.version
);
println!();
// Model line
let model_display = if let Some(ref cheap) = info.cheap_model {
format!(
"{cyan}{}{reset} {dim}cheap{reset} {cyan}{}{reset}",
info.llm_model, cheap
"{}{}{} {}cheap{} {}{}{}",
fmt::accent(),
info.llm_model,
fmt::reset(),
fmt::dim(),
fmt::reset(),
fmt::accent(),
cheap,
fmt::reset(),
)
} else {
format!("{cyan}{}{reset}", info.llm_model)
format!("{}{}{}", fmt::accent(), info.llm_model, fmt::reset())
};
println!(
" {dim}model{reset} {model_display} {dim}via {}{reset}",
info.llm_backend
" {}{:<width$}{} {model_display} {}via {}{}",
fmt::dim(),
"model",
fmt::reset(),
fmt::dim(),
info.llm_backend,
fmt::reset(),
width = KW,
);
// Database line
let db_status = if info.db_connected {
"connected"
} else {
"none"
};
println!(
" {dim}database{reset} {cyan}{}{reset} {dim}({db_status}){reset}",
info.db_backend
);
// ── Tier 2: conditional ───────────────────────────────────────────
// Tools line
println!(
" {dim}tools{reset} {cyan}{}{reset} {dim}registered{reset}",
info.tool_count
);
// Features line
let mut features = Vec::new();
if info.embeddings_enabled {
if let Some(ref provider) = info.embeddings_provider {
features.push(format!("embeddings ({provider})"));
} else {
features.push("embeddings".to_string());
}
}
if info.heartbeat_enabled {
let mins = info.heartbeat_interval_secs / 60;
features.push(format!("heartbeat ({mins}m)"));
}
match info.docker_status {
crate::sandbox::detect::DockerStatus::Available => {
features.push("sandbox".to_string());
}
crate::sandbox::detect::DockerStatus::NotInstalled => {
features.push(format!("{yellow}sandbox (docker not installed){reset}"));
}
crate::sandbox::detect::DockerStatus::NotRunning => {
features.push(format!("{yellow}sandbox (docker not running){reset}"));
}
crate::sandbox::detect::DockerStatus::Disabled => {
// Don't show sandbox when disabled
}
}
if info.claude_code_enabled {
features.push("claude-code".to_string());
}
if info.routines_enabled {
features.push("routines".to_string());
}
if info.skills_enabled {
features.push("skills".to_string());
}
if !features.is_empty() {
println!(
" {dim}features{reset} {cyan}{}{reset}",
features.join(" ")
);
}
// Channels line
if !info.channels.is_empty() {
println!(
" {dim}channels{reset} {cyan}{}{reset}",
info.channels.join(" ")
);
}
// Gateway URL (highlighted)
// Gateway URL
if let Some(ref url) = info.gateway_url {
println!();
println!(" {dim}gateway{reset} {yellow_underline}{url}{reset}");
println!(
" {}{:<width$}{} {}{}{}",
fmt::dim(),
"gateway",
fmt::reset(),
fmt::link(),
url,
fmt::reset(),
width = KW,
);
}
// Tunnel URL
@@ -142,15 +109,140 @@ pub fn print_boot_screen(info: &BootInfo) {
let provider_tag = info
.tunnel_provider
.as_deref()
.map(|p| format!(" {dim}({p}){reset}"))
.map(|p| format!(" {}({}){}", fmt::dim(), p, fmt::reset()))
.unwrap_or_default();
println!(" {dim}tunnel{reset} {yellow_underline}{url}{reset}{provider_tag}");
println!(
" {}{:<width$}{} {}{}{}{}",
fmt::dim(),
"tunnel",
fmt::reset(),
fmt::link(),
url,
fmt::reset(),
provider_tag,
width = KW,
);
}
// Non-default channels (skip if only the default set)
let non_default: Vec<&str> = info
.channels
.iter()
.filter(|c| !matches!(c.as_str(), "repl" | "gateway"))
.map(|c| c.as_str())
.collect();
if !non_default.is_empty() {
println!(
" {}{:<width$}{} {}{}{}",
fmt::dim(),
"channels",
fmt::reset(),
fmt::accent(),
non_default.join(" "),
fmt::reset(),
width = KW,
);
}
// ── Tier 3: compact feature tags ──────────────────────────────────
let mut tags: Vec<String> = Vec::new();
// Database
if info.db_connected {
tags.push(format!("db:{}", info.db_backend));
}
// Tool count
if info.tool_count > 0 {
tags.push(format!("tools:{}", info.tool_count));
}
// Routines
if info.routines_enabled {
tags.push("routines".to_string());
}
// Heartbeat with interval
if info.heartbeat_enabled {
let interval = if info.heartbeat_interval_secs >= 3600
&& info.heartbeat_interval_secs.is_multiple_of(3600)
{
format!("{}h", info.heartbeat_interval_secs / 3600)
} else if info.heartbeat_interval_secs >= 60
&& info.heartbeat_interval_secs.is_multiple_of(60)
{
format!("{}m", info.heartbeat_interval_secs / 60)
} else {
format!("{}s", info.heartbeat_interval_secs)
};
tags.push(format!("heartbeat:{interval}"));
}
// Skills
if info.skills_enabled {
tags.push("skills".to_string());
}
// Sandbox / Docker
if info.sandbox_enabled {
let suffix = match info.docker_status {
crate::sandbox::detect::DockerStatus::Available => "",
crate::sandbox::detect::DockerStatus::NotRunning => ":stopped",
_ => ":unavail",
};
tags.push(format!("sandbox{suffix}"));
}
// Embeddings
if info.embeddings_enabled {
if let Some(ref provider) = info.embeddings_provider {
tags.push(format!("embeddings:{provider}"));
} else {
tags.push("embeddings".to_string());
}
}
// Claude Code bridge
if info.claude_code_enabled {
tags.push("claude-code".to_string());
}
if !tags.is_empty() {
println!(
" {}{:<width$}{} {}",
fmt::dim(),
"features",
fmt::reset(),
tags.join(" "),
width = KW,
);
}
// ── Footer ────────────────────────────────────────────────────────
println!();
println!("{border}");
println!();
println!(" /help for commands, /quit to exit");
// Startup elapsed
if let Some(elapsed) = info.startup_elapsed {
let millis = elapsed.as_millis();
let elapsed_str = if millis < 1000 {
format!("{millis}ms")
} else {
let secs = elapsed.as_secs_f64();
format!("{secs:.1}s")
};
println!(" {}ready in {}{}", fmt::dim(), elapsed_str, fmt::reset());
}
// Hint to run `ironclaw status` for full details
println!(
" {}Run `ironclaw status` for full system details.{}",
fmt::hint(),
fmt::reset()
);
println!();
}
@@ -187,6 +279,7 @@ mod tests {
],
tunnel_url: Some("https://abc123.ngrok.io".to_string()),
tunnel_provider: Some("ngrok".to_string()),
startup_elapsed: None,
};
// Should not panic
print_boot_screen(&info);
@@ -216,6 +309,7 @@ mod tests {
channels: vec![],
tunnel_url: None,
tunnel_provider: None,
startup_elapsed: None,
};
// Should not panic
print_boot_screen(&info);
@@ -245,6 +339,7 @@ mod tests {
channels: vec!["repl".to_string()],
tunnel_url: None,
tunnel_provider: None,
startup_elapsed: None,
};
// Should not panic
print_boot_screen(&info);
+6
View File
@@ -333,6 +333,12 @@ pub enum StatusUpdate {
},
/// Suggested follow-up messages for the user.
Suggestions { suggestions: Vec<String> },
/// Per-turn token usage and cost summary (shown as subtle metadata).
TurnCost {
input_tokens: u64,
output_tokens: u64,
cost_usd: String,
},
}
impl StatusUpdate {
+338 -126
View File
@@ -20,6 +20,7 @@
use std::borrow::Cow;
use std::io::{self, IsTerminal, Write};
use std::sync::Arc;
use std::sync::Mutex;
use std::sync::atomic::{AtomicBool, Ordering};
use async_trait::async_trait;
@@ -40,6 +41,7 @@ use tokio_stream::wrappers::ReceiverStream;
use crate::agent::truncate_for_preview;
use crate::bootstrap::ironclaw_base_dir;
use crate::channels::{Channel, IncomingMessage, MessageStream, OutgoingResponse, StatusUpdate};
use crate::cli::fmt;
use crate::error::ChannelError;
/// Max characters for tool result previews in the terminal.
@@ -119,7 +121,7 @@ impl Hinter for ReplHelper {
impl Highlighter for ReplHelper {
fn highlight_hint<'h>(&self, hint: &'h str) -> Cow<'h, str> {
Cow::Owned(format!("\x1b[90m{hint}\x1b[0m"))
Cow::Owned(format!("{}{hint}{}", fmt::dim(), fmt::reset()))
}
}
@@ -143,55 +145,207 @@ impl ConditionalEventHandler for EscInterruptHandler {
}
}
/// Approval action chosen by the interactive selector.
#[derive(Clone, Copy)]
enum ApprovalAction {
Approve,
Always,
Deny,
}
impl std::fmt::Display for ApprovalAction {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Approve => write!(f, "Approve (y)"),
Self::Always => write!(f, "Always approve (a)"),
Self::Deny => write!(f, "Deny (n)"),
}
}
}
impl ApprovalAction {
fn as_input(self) -> &'static str {
match self {
Self::Approve => "y",
Self::Always => "a",
Self::Deny => "n",
}
}
}
/// Interactive approval selector using crossterm raw mode.
/// Returns the approval action string ("y", "a", or "n").
fn run_approval_selector(allow_always: bool) -> Option<&'static str> {
use crossterm::{
cursor,
event::{self, Event as CtEvent, KeyCode as CtKeyCode, KeyEventKind},
execute,
terminal::{self, ClearType},
};
let options: Vec<ApprovalAction> = if allow_always {
vec![
ApprovalAction::Approve,
ApprovalAction::Always,
ApprovalAction::Deny,
]
} else {
vec![ApprovalAction::Approve, ApprovalAction::Deny]
};
let num = options.len();
let mut sel: usize = 0;
// Total lines: options + hint line
let total_lines = (num + 1) as u16;
let render = |sel: usize| {
let mut w = io::stderr();
let pipe = format!("{}{}", fmt::accent(), fmt::reset());
for (i, opt) in options.iter().enumerate() {
if i == sel {
let _ = write!(w, " {pipe} {}● {opt}{}\r\n", fmt::bold(), fmt::reset());
} else {
let _ = write!(w, " {pipe} {}○ {opt}{}\r\n", fmt::dim(), fmt::reset());
}
}
let _ = write!(
w,
" {}└{} {}↑↓ enter to select{}\r\n",
fmt::accent(),
fmt::reset(),
fmt::dim(),
fmt::reset()
);
let _ = w.flush();
};
let _ = terminal::enable_raw_mode();
render(sel);
let result = loop {
let Ok(evt) = event::read() else { break None };
if let CtEvent::Key(key) = evt {
if key.kind != KeyEventKind::Press {
continue;
}
match key.code {
CtKeyCode::Up | CtKeyCode::Char('k') => {
sel = if sel == 0 { num - 1 } else { sel - 1 };
}
CtKeyCode::Down | CtKeyCode::Char('j') => {
sel = (sel + 1) % num;
}
CtKeyCode::Enter => break Some(options[sel].as_input()),
CtKeyCode::Char('y') | CtKeyCode::Char('Y') => break Some("y"),
CtKeyCode::Char('a') | CtKeyCode::Char('A') if allow_always => break Some("a"),
CtKeyCode::Char('n') | CtKeyCode::Char('N') => break Some("n"),
CtKeyCode::Esc => break None,
_ => continue,
}
// Redraw: move up, clear, render
let mut w = io::stderr();
let _ = execute!(w, cursor::MoveUp(total_lines));
let _ = execute!(w, terminal::Clear(ClearType::FromCursorDown));
render(sel);
}
};
let _ = terminal::disable_raw_mode();
// Overwrite selector with the confirmed choice
let mut w = io::stderr();
let _ = execute!(w, cursor::MoveUp(total_lines));
let _ = execute!(w, terminal::Clear(ClearType::FromCursorDown));
let (label, color) = if let Some(action) = result {
let l = options
.iter()
.find(|o| o.as_input() == action)
.unwrap_or(&options[0]);
let c = if action == "n" {
fmt::error()
} else {
fmt::success()
};
(l.to_string(), c)
} else {
(ApprovalAction::Deny.to_string(), fmt::error())
};
let _ = writeln!(
w,
" {}└{} {color}● {label}{}",
fmt::accent(),
fmt::reset(),
fmt::reset()
);
result
}
/// Build a termimad skin with our color scheme.
fn make_skin() -> MadSkin {
let mut skin = MadSkin::default();
skin.set_headers_fg(termimad::crossterm::style::Color::Yellow);
skin.bold.set_fg(termimad::crossterm::style::Color::White);
skin.italic
.set_fg(termimad::crossterm::style::Color::Magenta);
skin.inline_code
.set_fg(termimad::crossterm::style::Color::Green);
skin.code_block
.set_fg(termimad::crossterm::style::Color::Green);
skin.set_headers_fg(crossterm::style::Color::Yellow);
skin.bold.set_fg(crossterm::style::Color::White);
skin.italic.set_fg(crossterm::style::Color::Magenta);
skin.inline_code.set_fg(crossterm::style::Color::Green);
skin.code_block.set_fg(crossterm::style::Color::Green);
skin.code_block.left_margin = 2;
skin
}
/// Truncate a string to `max_chars` using character boundaries.
///
/// For strings longer than `max_chars`, shows the first half and last half
/// separated by `...` so both ends are visible.
fn smart_truncate(s: &str, max_chars: usize) -> Cow<'_, str> {
let char_count = s.chars().count();
if char_count <= max_chars {
return Cow::Borrowed(s);
}
// Account for the 3-char "..." separator
let budget = max_chars.saturating_sub(3);
let head_len = budget / 2;
let tail_len = budget - head_len;
let head: String = s.chars().take(head_len).collect();
let tail: String = s
.chars()
.skip(char_count.saturating_sub(tail_len))
.collect();
Cow::Owned(format!("{head}...{tail}"))
}
/// Format JSON params as `key: value` lines for the approval card.
fn format_json_params(params: &serde_json::Value, indent: &str) -> String {
let max_val_len = fmt::term_width().saturating_sub(8);
match params {
serde_json::Value::Object(map) => {
let mut lines = Vec::new();
for (key, value) in map {
let val_str = match value {
serde_json::Value::String(s) => {
let display = if s.len() > 120 { &s[..120] } else { s };
format!("\x1b[32m\"{display}\"\x1b[0m")
let display = smart_truncate(s, max_val_len);
format!("{}\"{display}\"{}", fmt::success(), fmt::reset())
}
other => {
let rendered = other.to_string();
if rendered.len() > 120 {
format!("{}...", &rendered[..120])
} else {
rendered
}
smart_truncate(&rendered, max_val_len).into_owned()
}
};
lines.push(format!("{indent}\x1b[36m{key}\x1b[0m: {val_str}"));
lines.push(format!(
"{indent}{}{key}{}: {val_str}",
fmt::accent(),
fmt::reset()
));
}
lines.join("\n")
}
other => {
let pretty = serde_json::to_string_pretty(other).unwrap_or_else(|_| other.to_string());
let truncated = if pretty.len() > 300 {
format!("{}...", &pretty[..300])
} else {
pretty
};
let truncated = smart_truncate(&pretty, 300);
truncated
.lines()
.map(|l| format!("{indent}\x1b[90m{l}\x1b[0m"))
.map(|l| format!("{indent}{}{l}{}", fmt::dim(), fmt::reset()))
.collect::<Vec<_>>()
.join("\n")
}
@@ -210,6 +364,12 @@ pub struct ReplChannel {
is_streaming: Arc<AtomicBool>,
/// When true, the one-liner startup banner is suppressed (boot screen shown instead).
suppress_banner: Arc<AtomicBool>,
/// Sender to inject messages into the agent loop (set after start()).
msg_tx: Arc<Mutex<Option<mpsc::Sender<IncomingMessage>>>>,
/// When true, the readline thread must yield stdin (approval selector or agent processing).
stdin_locked: Arc<AtomicBool>,
/// Number of transient status lines (Thinking) to erase on next output.
transient_lines: std::sync::atomic::AtomicU8,
}
impl ReplChannel {
@@ -226,6 +386,9 @@ impl ReplChannel {
debug_mode: Arc::new(AtomicBool::new(false)),
is_streaming: Arc::new(AtomicBool::new(false)),
suppress_banner: Arc::new(AtomicBool::new(false)),
msg_tx: Arc::new(Mutex::new(None)),
stdin_locked: Arc::new(AtomicBool::new(false)),
transient_lines: std::sync::atomic::AtomicU8::new(0),
}
}
@@ -242,6 +405,9 @@ impl ReplChannel {
debug_mode: Arc::new(AtomicBool::new(false)),
is_streaming: Arc::new(AtomicBool::new(false)),
suppress_banner: Arc::new(AtomicBool::new(false)),
msg_tx: Arc::new(Mutex::new(None)),
stdin_locked: Arc::new(AtomicBool::new(false)),
transient_lines: std::sync::atomic::AtomicU8::new(0),
}
}
@@ -253,6 +419,17 @@ impl ReplChannel {
fn is_debug(&self) -> bool {
self.debug_mode.load(Ordering::Relaxed)
}
/// Erase transient status lines (Thinking indicators) from the terminal.
fn clear_transient(&self) {
use crossterm::{cursor, execute, terminal};
let n = self.transient_lines.swap(0, Ordering::Relaxed);
if n > 0 {
let mut stderr = io::stderr();
let _ = execute!(stderr, cursor::MoveUp(n as u16));
let _ = execute!(stderr, terminal::Clear(terminal::ClearType::FromCursorDown));
}
}
}
impl Default for ReplChannel {
@@ -262,33 +439,30 @@ impl Default for ReplChannel {
}
fn print_help() {
// Bold white for section headers, bold cyan for commands, dim gray for descriptions
let h = "\x1b[1m"; // bold (section headers)
let c = "\x1b[1;36m"; // bold cyan (commands)
let d = "\x1b[90m"; // dim gray (descriptions)
let r = "\x1b[0m"; // reset
let h = fmt::bold();
let c = fmt::bold_accent();
let d = fmt::dim();
let r = fmt::reset();
let hi = fmt::hint();
println!();
println!(" {h}IronClaw REPL{r}");
println!();
println!(" {h}Commands{r}");
println!(" {c}/help{r} {d}show this help{r}");
println!(" {c}/debug{r} {d}toggle verbose output{r}");
println!(" {c}/quit{r} {c}/exit{r} {d}exit the repl{r}");
println!(" {h}Quick start{r}");
println!(" {c}/new{r} {hi}Start a new thread{r}");
println!(" {c}/compact{r} {hi}Compress context window{r}");
println!(" {c}/quit{r} {hi}Exit{r}");
println!();
println!(" {h}Conversation{r}");
println!(" {c}/undo{r} {d}undo the last turn{r}");
println!(" {c}/redo{r} {d}redo an undone turn{r}");
println!(" {c}/clear{r} {d}clear conversation{r}");
println!(" {c}/compact{r} {d}compact context window{r}");
println!(" {c}/new{r} {d}new conversation thread{r}");
println!(" {c}/interrupt{r} {d}stop current operation{r}");
println!(" {c}esc{r} {d}stop current operation{r}");
println!();
println!(" {h}Approval responses{r}");
println!(" {c}yes{r} ({c}y{r}) {d}approve tool execution{r}");
println!(" {c}no{r} ({c}n{r}) {d}deny tool execution{r}");
println!(" {c}always{r} ({c}a{r}) {d}approve for this session{r}");
println!(" {h}All commands{r}");
println!(
" {d}Conversation{r} {c}/new{r} {c}/clear{r} {c}/compact{r} {c}/undo{r} {c}/redo{r} {c}/summarize{r} {c}/suggest{r}"
);
println!(" {d}Threads{r} {c}/thread{r} {c}/resume{r} {c}/list{r}");
println!(" {d}Execution{r} {c}/interrupt{r} {d}(esc){r} {c}/cancel{r}");
println!(
" {d}System{r} {c}/tools{r} {c}/model{r} {c}/version{r} {c}/status{r} {c}/debug{r} {c}/heartbeat{r}"
);
println!(" {d}Session{r} {c}/help{r} {c}/quit{r}");
println!();
}
@@ -305,10 +479,15 @@ impl Channel for ReplChannel {
async fn start(&self) -> Result<MessageStream, ChannelError> {
let (tx, rx) = mpsc::channel(32);
// Store tx so send_status can inject approval responses directly
if let Ok(mut guard) = self.msg_tx.lock() {
*guard = Some(tx.clone());
}
let single_message = self.single_message.clone();
let user_id = self.user_id.clone();
let debug_mode = Arc::clone(&self.debug_mode);
let suppress_banner = Arc::clone(&self.suppress_banner);
let stdin_locked = Arc::clone(&self.stdin_locked);
let esc_interrupt_triggered_for_thread = Arc::new(AtomicBool::new(false));
std::thread::spawn(move || {
@@ -357,18 +536,33 @@ impl Channel for ReplChannel {
let _ = rl.load_history(&hist_path);
if !suppress_banner.load(Ordering::Relaxed) {
println!("\x1b[1mIronClaw\x1b[0m /help for commands, /quit to exit");
println!(
"{}IronClaw{} /help for commands, /quit to exit",
fmt::bold(),
fmt::reset()
);
println!();
}
loop {
// Yield stdin while approval selector or agent processing locks it
while stdin_locked.load(Ordering::Relaxed) {
std::thread::sleep(std::time::Duration::from_millis(50));
}
let prompt = if debug_mode.load(Ordering::Relaxed) {
"\x1b[33m[debug]\x1b[0m \x1b[1;36m\u{203A}\x1b[0m "
format!(
"{}[debug]{} {}\u{203A}{} ",
fmt::warning(),
fmt::reset(),
fmt::bold_accent(),
fmt::reset()
)
} else {
"\x1b[1;36m\u{203A}\x1b[0m "
format!("{}\u{203A}{} ", fmt::bold_accent(), fmt::reset())
};
match rl.readline(prompt) {
match rl.readline(&prompt) {
Ok(line) => {
let line = line.trim();
if line.is_empty() {
@@ -394,9 +588,9 @@ impl Channel for ReplChannel {
let current = debug_mode.load(Ordering::Relaxed);
debug_mode.store(!current, Ordering::Relaxed);
if !current {
println!("\x1b[90mdebug mode on\x1b[0m");
println!("{}debug mode on{}", fmt::dim(), fmt::reset());
} else {
println!("\x1b[90mdebug mode off\x1b[0m");
println!("{}debug mode off{}", fmt::dim(), fmt::reset());
}
continue;
}
@@ -405,7 +599,11 @@ impl Channel for ReplChannel {
let msg =
IncomingMessage::new("repl", &user_id, line).with_timezone(&sys_tz);
// Lock stdin before sending so readline doesn't restart
// while the agent is processing (approval selector needs stdin)
stdin_locked.store(true, Ordering::Relaxed);
if tx.blocking_send(msg).is_err() {
stdin_locked.store(false, Ordering::Relaxed);
break;
}
}
@@ -456,21 +654,23 @@ impl Channel for ReplChannel {
_msg: &IncomingMessage,
response: OutgoingResponse,
) -> Result<(), ChannelError> {
let width = crossterm::terminal::size()
.map(|(w, _)| w as usize)
.unwrap_or(80);
let width = fmt::term_width();
// If we were streaming, the content was already printed via StreamChunk.
// Just finish the line and reset.
if self.is_streaming.swap(false, Ordering::Relaxed) {
println!();
println!();
self.stdin_locked.store(false, Ordering::Relaxed);
return Ok(());
}
// Clear any leftover thinking indicators
self.clear_transient();
// Dim separator line before the response
let sep_width = width.min(80);
eprintln!("\x1b[90m{}\x1b[0m", "\u{2500}".repeat(sep_width));
eprintln!("{}", fmt::separator(sep_width));
// Render markdown
let skin = make_skin();
@@ -478,6 +678,8 @@ impl Channel for ReplChannel {
print!("{text}");
println!();
// Unlock stdin so readline can resume
self.stdin_locked.store(false, Ordering::Relaxed);
Ok(())
}
@@ -490,31 +692,34 @@ impl Channel for ReplChannel {
match status {
StatusUpdate::Thinking(msg) => {
self.clear_transient();
let display = truncate_for_preview(&msg, CLI_STATUS_MAX);
eprintln!(" \x1b[90m\u{25CB} {display}\x1b[0m");
eprintln!(" {}\u{25CB} {display}{}", fmt::dim(), fmt::reset());
self.transient_lines.store(1, Ordering::Relaxed);
}
StatusUpdate::ToolStarted { name } => {
eprintln!(" \x1b[33m\u{25CB} {name}\x1b[0m");
self.clear_transient();
eprintln!(" {}\u{25CB} {name}{}", fmt::dim(), fmt::reset());
self.transient_lines.store(1, Ordering::Relaxed);
}
StatusUpdate::ToolCompleted { name, success, .. } => {
self.clear_transient();
if success {
eprintln!(" \x1b[32m\u{25CF} {name}\x1b[0m");
eprintln!(" {}\u{25CF} {name}{}", fmt::success(), fmt::reset());
} else {
eprintln!(" \x1b[31m\u{2717} {name} (failed)\x1b[0m");
eprintln!(" {}\u{2717} {name} (failed){}", fmt::error(), fmt::reset());
}
}
StatusUpdate::ToolResult { name: _, preview } => {
let display = truncate_for_preview(&preview, CLI_TOOL_RESULT_MAX);
eprintln!(" \x1b[90m{display}\x1b[0m");
eprintln!(" {}{display}{}", fmt::dim(), fmt::reset());
}
StatusUpdate::StreamChunk(chunk) => {
// Print separator on the false-to-true transition
if !self.is_streaming.swap(true, Ordering::Relaxed) {
let width = crossterm::terminal::size()
.map(|(w, _)| w as usize)
.unwrap_or(80);
let sep_width = width.min(80);
eprintln!("\x1b[90m{}\x1b[0m", "\u{2500}".repeat(sep_width));
self.clear_transient();
let sep_width = fmt::term_width().min(80);
eprintln!("{}", fmt::separator(sep_width));
}
print!("{chunk}");
let _ = io::stdout().flush();
@@ -525,73 +730,67 @@ impl Channel for ReplChannel {
browse_url,
} => {
eprintln!(
" \x1b[36m[job]\x1b[0m {title} \x1b[90m({job_id})\x1b[0m \x1b[4m{browse_url}\x1b[0m"
" {}[job]{} {title} {}({job_id}){} {}{browse_url}{}",
fmt::accent(),
fmt::reset(),
fmt::dim(),
fmt::reset(),
fmt::link(),
fmt::reset()
);
}
StatusUpdate::Status(msg) => {
if debug || msg.contains("approval") || msg.contains("Approval") {
let display = truncate_for_preview(&msg, CLI_STATUS_MAX);
eprintln!(" \x1b[90m{display}\x1b[0m");
eprintln!(" {}{display}{}", fmt::dim(), fmt::reset());
}
}
StatusUpdate::ApprovalNeeded {
request_id,
request_id: _,
tool_name,
description,
description: _,
parameters,
allow_always,
} => {
let term_width = crossterm::terminal::size()
.map(|(w, _)| w as usize)
.unwrap_or(80);
let box_width = (term_width.saturating_sub(4)).clamp(40, 60);
self.clear_transient();
let pipe = format!("{}{}", fmt::accent(), fmt::reset());
// Short request ID for the bottom border
let short_id = if request_id.len() > 8 {
&request_id[..8]
} else {
&request_id
};
// Top border: ┌ tool_name requires approval ───
let top_label = format!(" {tool_name} requires approval ");
let top_fill = box_width.saturating_sub(top_label.len() + 1);
let top_border = format!(
"\u{250C}\x1b[33m{top_label}\x1b[0m{}",
"\u{2500}".repeat(top_fill)
// Header: ◆ tool requires approval
eprintln!();
eprintln!(
" {}\u{25C6} {}{tool_name}{} requires approval",
fmt::accent(),
fmt::bold(),
fmt::reset()
);
// Bottom border: └─ short_id ─────
let bot_label = format!(" {short_id} ");
let bot_fill = box_width.saturating_sub(bot_label.len() + 2);
let bot_border = format!(
"\u{2514}\u{2500}\x1b[90m{bot_label}\x1b[0m{}",
"\u{2500}".repeat(bot_fill)
);
eprintln!();
eprintln!(" {top_border}");
eprintln!(" \u{2502} \x1b[90m{description}\x1b[0m");
eprintln!(" \u{2502}");
// Params
let param_lines = format_json_params(&parameters, " \u{2502} ");
// The format_json_params already includes the indent prefix
// but we need to handle the case where each line already starts with it
for line in param_lines.lines() {
eprintln!("{line}");
// Params: │ key value
let param_lines = format_json_params(&parameters, &format!(" {pipe} "));
if !param_lines.is_empty() {
eprintln!(" {pipe}");
for line in param_lines.lines() {
eprintln!("{line}");
}
}
eprintln!(" \u{2502}");
if allow_always {
eprintln!(
" \u{2502} \x1b[32myes\x1b[0m (y) / \x1b[34malways\x1b[0m (a) / \x1b[31mno\x1b[0m (n)"
);
} else {
eprintln!(" \u{2502} \x1b[32myes\x1b[0m (y) / \x1b[31mno\x1b[0m (n)");
}
eprintln!(" {bot_border}");
eprintln!();
eprintln!(" {pipe}");
// Run interactive selector directly from send_status
// stdin is already locked by Thinking/ToolStarted, so the
// readline thread is not competing for stdin.
let msg_tx = Arc::clone(&self.msg_tx);
let user_id = self.user_id.clone();
let lock_flag = Arc::clone(&self.stdin_locked);
tokio::task::spawn_blocking(move || {
let action = run_approval_selector(allow_always).unwrap_or("n");
// Unlock stdin so readline can resume after approval
lock_flag.store(false, Ordering::Relaxed);
let Ok(guard) = msg_tx.lock() else {
return;
};
if let Some(tx) = guard.as_ref() {
let msg = IncomingMessage::new("repl", &user_id, action);
let _ = tx.blocking_send(msg);
}
});
}
StatusUpdate::AuthRequired {
extension_name,
@@ -600,12 +799,16 @@ impl Channel for ReplChannel {
..
} => {
eprintln!();
eprintln!("\x1b[33m Authentication required for {extension_name}\x1b[0m");
eprintln!(
"{} Authentication required for {extension_name}{}",
fmt::warning(),
fmt::reset()
);
if let Some(ref instr) = instructions {
eprintln!(" {instr}");
}
if let Some(ref url) = setup_url {
eprintln!(" \x1b[4m{url}\x1b[0m");
eprintln!(" {}{url}{}", fmt::link(), fmt::reset());
}
eprintln!();
}
@@ -615,21 +818,32 @@ impl Channel for ReplChannel {
message,
} => {
if success {
eprintln!("\x1b[32m {extension_name}: {message}\x1b[0m");
eprintln!(
"{} {extension_name}: {message}{}",
fmt::success(),
fmt::reset()
);
} else {
eprintln!("\x1b[31m {extension_name}: {message}\x1b[0m");
eprintln!(
"{} {extension_name}: {message}{}",
fmt::error(),
fmt::reset()
);
}
}
StatusUpdate::ImageGenerated { path, .. } => {
if let Some(ref p) = path {
eprintln!("\x1b[36m [image] {p}\x1b[0m");
eprintln!("{} [image] {p}{}", fmt::accent(), fmt::reset());
} else {
eprintln!("\x1b[36m [image generated]\x1b[0m");
eprintln!("{} [image generated]{}", fmt::accent(), fmt::reset());
}
}
StatusUpdate::Suggestions { .. } => {
// Suggestions are only rendered by the web gateway
}
StatusUpdate::TurnCost { .. } => {
// Cost display is handled by the TUI channel
}
}
Ok(())
}
@@ -640,11 +854,9 @@ impl Channel for ReplChannel {
response: OutgoingResponse,
) -> Result<(), ChannelError> {
let skin = make_skin();
let width = crossterm::terminal::size()
.map(|(w, _)| w as usize)
.unwrap_or(80);
let width = fmt::term_width();
eprintln!("\x1b[34m\u{25CF}\x1b[0m notification");
eprintln!("{}\u{25CF}{} notification", fmt::accent(), fmt::reset());
let text = termimad::FmtText::from(&skin, &response.content, Some(width));
eprint!("{text}");
eprintln!();
+1 -1
View File
@@ -117,7 +117,7 @@ async fn register_channel(
wasm_router: &Arc<WasmChannelRouter>,
) -> (String, Box<dyn crate::channels::Channel>) {
let channel_name = loaded.name().to_string();
tracing::info!("Loaded WASM channel: {}", channel_name);
tracing::debug!("Loaded WASM channel: {}", channel_name);
let owner_actor_id = config
.channels
.wasm_channel_owner_ids
+2 -2
View File
@@ -3059,8 +3059,8 @@ fn status_to_wit(
},
metadata_json,
},
// Suggestions are web-gateway-only; skip for WASM channels
StatusUpdate::Suggestions { .. } => return None,
// Suggestions and turn cost are web-gateway-only; skip for WASM channels
StatusUpdate::Suggestions { .. } | StatusUpdate::TurnCost { .. } => return None,
})
}
+10
View File
@@ -415,6 +415,16 @@ impl Channel for GatewayChannel {
suggestions,
thread_id,
},
StatusUpdate::TurnCost {
input_tokens,
output_tokens,
cost_usd,
} => SseEvent::TurnCost {
input_tokens,
output_tokens,
cost_usd,
thread_id,
},
};
self.state.sse.broadcast(event);
+7 -3
View File
@@ -2343,7 +2343,7 @@ async fn extensions_setup_handler(
"Extension manager not available (secrets store required)".to_string(),
))?;
let secrets = ext_mgr
let setup = ext_mgr
.get_setup_schema(&name)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
@@ -2359,7 +2359,8 @@ async fn extensions_setup_handler(
Ok(Json(ExtensionSetupResponse {
name,
kind,
secrets,
secrets: setup.secrets,
fields: setup.fields,
}))
}
@@ -2377,7 +2378,7 @@ async fn extensions_setup_submit_handler(
// through to the LLM instead of being intercepted as a token.
clear_auth_mode(&state).await;
match ext_mgr.configure(&name, &req.secrets).await {
match ext_mgr.configure(&name, &req.secrets, &req.fields).await {
Ok(result) => {
let mut resp = if result.verification.is_some() || result.activated {
ActionResponse::ok(result.message)
@@ -2385,6 +2386,9 @@ async fn extensions_setup_submit_handler(
ActionResponse::fail(result.message)
};
resp.activated = Some(result.activated);
if result.restart_required || !result.activated {
resp.needs_restart = Some(true);
}
resp.auth_url = result.auth_url.clone();
resp.verification = result.verification.clone();
resp.instructions = result.verification.as_ref().map(|v| v.instructions.clone());
+1
View File
@@ -144,6 +144,7 @@ impl SseManager {
SseEvent::Heartbeat => "heartbeat",
SseEvent::ImageGenerated { .. } => "image_generated",
SseEvent::Suggestions { .. } => "suggestions",
SseEvent::TurnCost { .. } => "turn_cost",
SseEvent::ExtensionStatus { .. } => "extension_status",
};
Ok(Event::default().event(event_type).data(data))
File diff suppressed because it is too large Load Diff
+25
View File
@@ -521,4 +521,29 @@ I18n.register('en', {
'channels.replDesc': 'Simple read-eval-print loop for testing',
'channels.configureVia': 'Configure via {env}',
'channels.runWith': 'Run with: {cmd}',
// Welcome Card
'welcome.heading': 'What can I help you with?',
'welcome.description': 'IronClaw is your secure AI assistant. Choose a suggestion below or type your own message.',
'welcome.runTool': 'Run a tool',
'welcome.checkJobs': 'Check job status',
'welcome.searchMemory': 'Search memory',
'welcome.manageRoutines': 'Manage routines',
'welcome.systemStatus': 'System status',
'welcome.writeCode': 'Write code',
// Connection
'connection.disconnected': 'Disconnected — attempting to reconnect',
'connection.reconnecting': 'Reconnecting (attempt {count})...',
'connection.reconnected': 'Reconnected',
// Messages
'message.you': 'You',
'message.assistant': 'IronClaw',
'message.system': 'System',
'message.copy': 'Copy',
'message.copied': 'Copied!',
// Approval
'approval.pressY': 'Press Y to approve, N to deny',
});
+25
View File
@@ -520,4 +520,29 @@ I18n.register('zh-CN', {
'channels.replDesc': '用于测试的简单读取-求值-打印循环',
'channels.configureVia': '通过 {env} 配置',
'channels.runWith': '运行命令: {cmd}',
// Welcome Card
'welcome.heading': '有什么可以帮助您的?',
'welcome.description': 'IronClaw 是您的安全 AI 助手。选择下方的建议或输入您自己的消息。',
'welcome.runTool': '运行工具',
'welcome.checkJobs': '查看任务状态',
'welcome.searchMemory': '搜索记忆',
'welcome.manageRoutines': '管理例程',
'welcome.systemStatus': '系统状态',
'welcome.writeCode': '编写代码',
// Connection
'connection.disconnected': '已断开连接 — 正在尝试重新连接',
'connection.reconnecting': '正在重新连接(第 {count} 次尝试)...',
'connection.reconnected': '已重新连接',
// Messages
'message.you': '你',
'message.assistant': 'IronClaw',
'message.system': '系统',
'message.copy': '复制',
'message.copied': '已复制!',
// Approval
'approval.pressY': '按 Y 批准,N 拒绝',
});
+3
View File
@@ -92,6 +92,7 @@
<div id="app">
<!-- Tab Bar -->
<div class="tab-bar">
<div class="tab-indicator" id="tab-indicator"></div>
<button class="active" data-tab="chat" data-i18n="tab.chat">Chat</button>
<button data-tab="memory" data-i18n="tab.memory">Memory</button>
<button data-tab="jobs" data-i18n="tab.jobs">Jobs</button>
@@ -292,9 +293,11 @@
<button class="settings-subtab" data-settings-subtab="extensions" data-i18n="tab.extensions">Extensions</button>
<button class="settings-subtab" data-settings-subtab="mcp" data-i18n="settings.mcp">MCP</button>
<button class="settings-subtab" data-settings-subtab="skills" data-i18n="tab.skills">Skills</button>
<button class="settings-theme-toggle" id="settings-theme-toggle" data-i18n="theme.tooltipSystem" title="Toggle theme">Theme</button>
</div>
<div class="settings-content">
<div class="settings-toolbar">
<button id="settings-back-btn" class="settings-back-btn">&larr; Back</button>
<div class="settings-search">
<input type="text" id="settings-search-input" data-i18n-placeholder="settings.searchPlaceholder" placeholder="Search settings..." data-i18n-attr="aria-label" data-i18n="settings.searchPlaceholder" aria-label="Search settings...">
</div>
File diff suppressed because it is too large Load Diff
+65
View File
@@ -254,6 +254,16 @@ pub enum SseEvent {
thread_id: Option<String>,
},
/// Per-turn token usage and cost summary.
#[serde(rename = "turn_cost")]
TurnCost {
input_tokens: u64,
output_tokens: u64,
cost_usd: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
/// Extension activation status change (WASM channels).
#[serde(rename = "extension_status")]
ExtensionStatus {
@@ -525,6 +535,7 @@ pub struct ExtensionSetupResponse {
pub name: String,
pub kind: String,
pub secrets: Vec<SecretFieldInfo>,
pub fields: Vec<SetupFieldInfo>,
}
#[derive(Debug, Serialize)]
@@ -538,9 +549,23 @@ pub struct SecretFieldInfo {
pub auto_generate: bool,
}
#[derive(Debug, Serialize)]
pub struct SetupFieldInfo {
pub name: String,
pub prompt: String,
pub optional: bool,
/// Whether this field already has a stored value.
pub provided: bool,
/// Input type for web UI rendering.
pub input_type: crate::tools::wasm::ToolSetupFieldInputType,
}
#[derive(Debug, Deserialize)]
pub struct ExtensionSetupRequest {
#[serde(default)]
pub secrets: std::collections::HashMap<String, String>,
#[serde(default)]
pub fields: std::collections::HashMap<String, String>,
}
#[derive(Debug, Serialize)]
@@ -559,6 +584,9 @@ pub struct ActionResponse {
/// Whether the channel was successfully activated after setup.
#[serde(skip_serializing_if = "Option::is_none")]
pub activated: Option<bool>,
/// Whether a restart is required for the new configuration to take effect.
#[serde(skip_serializing_if = "Option::is_none")]
pub needs_restart: Option<bool>,
/// Pending manual verification challenge (for Telegram owner binding, etc.).
#[serde(skip_serializing_if = "Option::is_none")]
pub verification: Option<crate::extensions::VerificationChallenge>,
@@ -573,6 +601,7 @@ impl ActionResponse {
awaiting_token: None,
instructions: None,
activated: None,
needs_restart: None,
verification: None,
}
}
@@ -585,6 +614,7 @@ impl ActionResponse {
awaiting_token: None,
instructions: None,
activated: None,
needs_restart: None,
verification: None,
}
}
@@ -777,6 +807,7 @@ impl WsServerMessage {
SseEvent::JobResult { .. } => "job_result",
SseEvent::ImageGenerated { .. } => "image_generated",
SseEvent::Suggestions { .. } => "suggestions",
SseEvent::TurnCost { .. } => "turn_cost",
SseEvent::ExtensionStatus { .. } => "extension_status",
};
let data = serde_json::to_value(event).unwrap_or(serde_json::Value::Null);
@@ -1246,6 +1277,40 @@ mod tests {
assert_eq!(req.extension_name, "telegram");
}
#[test]
fn test_extension_setup_request_defaults() {
let json = r#"{}"#;
let req: ExtensionSetupRequest = serde_json::from_str(json).unwrap();
assert!(req.secrets.is_empty());
assert!(req.fields.is_empty());
}
#[test]
fn test_extension_setup_request_deserialize_with_fields() {
let json = r#"{
"secrets": { "api_key": "sk-123" },
"fields": { "llm_backend": "openai", "selected_model": "gpt-4o" }
}"#;
let req: ExtensionSetupRequest = serde_json::from_str(json).unwrap();
assert_eq!(req.secrets.get("api_key").unwrap(), "sk-123");
assert_eq!(req.fields.get("llm_backend").unwrap(), "openai");
assert_eq!(req.fields.get("selected_model").unwrap(), "gpt-4o");
}
#[test]
fn test_setup_field_info_serializes_input_type_as_enum_string() {
let field = SetupFieldInfo {
name: "selected_model".to_string(),
prompt: "Model".to_string(),
optional: false,
provided: true,
input_type: crate::tools::wasm::ToolSetupFieldInputType::Password,
};
let json = serde_json::to_value(field).unwrap();
assert_eq!(json["input_type"], "password");
}
// ---- ThreadInfo channel field tests ----
#[test]
+2 -2
View File
@@ -175,7 +175,7 @@ mod tests {
#[test]
fn test_truncate_preview_closes_tool_output_tag() {
let s = "<tool_output name=\"search\" sanitized=\"true\">\nSome very long content here\n</tool_output>";
let s = "<tool_output name=\"search\">\nSome very long content here\n</tool_output>";
// Truncate so it cuts before the closing tag
let result = truncate_preview(s, 60);
assert!(result.ends_with("</tool_output>"));
@@ -184,7 +184,7 @@ mod tests {
#[test]
fn test_truncate_preview_no_extra_close_when_intact() {
let s = "<tool_output name=\"echo\" sanitized=\"false\">\nshort\n</tool_output>";
let s = "<tool_output name=\"echo\">\nshort\n</tool_output>";
// The string is short enough not to be truncated
let result = truncate_preview(s, 500);
assert_eq!(result, s);
+2 -2
View File
@@ -68,7 +68,7 @@ impl WebhookServer {
reason: format!("Failed to bind to {}: {}", self.config.addr, e),
})?;
tracing::info!("Webhook server listening on {}", self.config.addr);
tracing::debug!("Webhook server listening on {}", self.config.addr);
let (shutdown_tx, shutdown_rx) = oneshot::channel();
self.shutdown_tx = Some(shutdown_tx);
@@ -129,7 +129,7 @@ impl WebhookServer {
});
self.handle = Some(handle);
tracing::info!("Webhook server listening on {}", new_addr);
tracing::debug!("Webhook server listening on {}", new_addr);
(old_shutdown_tx, old_handle)
}
+44 -9
View File
@@ -7,12 +7,13 @@
use std::path::PathBuf;
use crate::bootstrap::ironclaw_base_dir;
use crate::cli::fmt;
use crate::settings::Settings;
/// Run all diagnostic checks and print results.
pub async fn run_doctor_command() -> anyhow::Result<()> {
println!("IronClaw Doctor");
println!("===============\n");
println!();
println!(" {}IronClaw Doctor{}", fmt::bold(), fmt::reset());
let mut passed = 0u32;
let mut failed = 0u32;
@@ -21,7 +22,9 @@ pub async fn run_doctor_command() -> anyhow::Result<()> {
// Load settings once for checks that need them.
let settings = Settings::load();
// ── Settings & core config ─────────────────────────────────
// ── Core ─────────────────────────────────────────────────
section_header("Core");
check(
"Settings file",
@@ -63,7 +66,9 @@ pub async fn run_doctor_command() -> anyhow::Result<()> {
&mut skipped,
);
// ── Subsystem configuration checks ─────────────────────────
// ── Features ─────────────────────────────────────────────
section_header("Features");
check(
"Embeddings",
@@ -121,7 +126,9 @@ pub async fn run_doctor_command() -> anyhow::Result<()> {
&mut skipped,
);
// ── External binary checks ────────────────────────────────
// ── External ─────────────────────────────────────────────
section_header("External");
check(
"Docker daemon",
@@ -158,7 +165,18 @@ pub async fn run_doctor_command() -> anyhow::Result<()> {
// ── Summary ───────────────────────────────────────────────
println!();
println!(" {passed} passed, {failed} failed, {skipped} skipped");
println!(
" {}{} passed{}, {}{} failed{}, {}{} skipped{}",
fmt::success(),
passed,
fmt::reset(),
if failed > 0 { fmt::error() } else { fmt::dim() },
failed,
fmt::reset(),
fmt::dim(),
skipped,
fmt::reset(),
);
if failed > 0 {
println!("\n Some checks failed. This is normal if you don't use those features.");
@@ -167,21 +185,38 @@ pub async fn run_doctor_command() -> anyhow::Result<()> {
Ok(())
}
/// Print a section header with a separator and bold group name.
fn section_header(name: &str) {
println!();
println!(" {}", fmt::separator(36));
println!(" {}{}{}", fmt::bold(), name, fmt::reset());
println!();
}
// ── Individual checks ───────────────────────────────────────
fn check(name: &str, result: CheckResult, passed: &mut u32, failed: &mut u32, skipped: &mut u32) {
match result {
CheckResult::Pass(detail) => {
*passed += 1;
println!(" [pass] {name}: {detail}");
println!(
"{}",
fmt::check_line(fmt::StatusKind::Pass, name, &detail, 18)
);
}
CheckResult::Fail(detail) => {
*failed += 1;
println!(" [FAIL] {name}: {detail}");
println!(
"{}",
fmt::check_line(fmt::StatusKind::Fail, name, &detail, 18)
);
}
CheckResult::Skip(reason) => {
*skipped += 1;
println!(" [skip] {name}: {reason}");
println!(
"{}",
fmt::check_line(fmt::StatusKind::Skip, name, &reason, 18)
);
}
}
}
+296
View File
@@ -0,0 +1,296 @@
//! Shared terminal design system.
//!
//! Centralizes color tokens, rendering primitives, and width detection
//! for consistent CLI output. Respects `NO_COLOR` env var and non-TTY
//! output (piping to file, CI, etc.).
use std::io::IsTerminal;
// ── Color detection ─────────────────────────────────────────
/// Returns `true` when ANSI colors should be emitted.
///
/// Disabled when:
/// - `NO_COLOR` env var is set (any value — per <https://no-color.org/>)
/// - stdout is not a terminal (pipe, file redirect, CI)
fn colors_enabled() -> bool {
if std::env::var_os("NO_COLOR").is_some() {
return false;
}
std::io::stdout().is_terminal()
}
/// Returns `true` when the terminal supports 24-bit true-color.
///
/// Checks `$COLORTERM` for `truecolor` or `24bit`.
fn truecolor_enabled() -> bool {
std::env::var("COLORTERM")
.map(|v| v.eq_ignore_ascii_case("truecolor") || v.eq_ignore_ascii_case("24bit"))
.unwrap_or(false)
}
// ── Color tokens ────────────────────────────────────────────
/// Emerald green accent — primary brand color.
///
/// Uses true-color `#34d399` when supported, falls back to basic green.
pub fn accent() -> &'static str {
if !colors_enabled() {
return "";
}
if truecolor_enabled() {
"\x1b[38;2;52;211;153m"
} else {
"\x1b[32m"
}
}
/// Bold text.
pub fn bold() -> &'static str {
if colors_enabled() { "\x1b[1m" } else { "" }
}
/// Green — success indicators.
pub fn success() -> &'static str {
if colors_enabled() { "\x1b[32m" } else { "" }
}
/// Yellow — warning indicators.
pub fn warning() -> &'static str {
if colors_enabled() { "\x1b[33m" } else { "" }
}
/// Red — error indicators.
pub fn error() -> &'static str {
if colors_enabled() { "\x1b[31m" } else { "" }
}
/// Dim gray — labels, secondary text.
pub fn dim() -> &'static str {
if colors_enabled() { "\x1b[90m" } else { "" }
}
/// Yellow underline — URLs and links.
pub fn link() -> &'static str {
if colors_enabled() { "\x1b[33;4m" } else { "" }
}
/// Bold accent — commands and interactive elements.
///
/// Uses bold + true-color emerald when supported, falls back to bold green.
pub fn bold_accent() -> &'static str {
if !colors_enabled() {
return "";
}
if truecolor_enabled() {
"\x1b[1;38;2;52;211;153m"
} else {
"\x1b[1;32m"
}
}
/// Dim italic — contextual tips and hints.
pub fn hint() -> &'static str {
if colors_enabled() { "\x1b[2;3m" } else { "" }
}
/// Reset all attributes.
pub fn reset() -> &'static str {
if colors_enabled() { "\x1b[0m" } else { "" }
}
// ── Width detection ─────────────────────────────────────────
/// Detect terminal width, clamped to [40, 120].
pub fn term_width() -> usize {
crossterm::terminal::size()
.map(|(w, _)| w as usize)
.unwrap_or(80)
.clamp(40, 120)
}
// ── Rendering primitives ────────────────────────────────────
/// Horizontal separator line (dim `─` characters).
pub fn separator(width: usize) -> String {
format!("{}{}{}", dim(), "\u{2500}".repeat(width), reset())
}
/// Key-value line with right-padded dim key and accent value.
///
/// ```text
/// Database libsql (connected)
/// ```
pub fn kv_line(key: &str, value: &str, key_width: usize) -> String {
format!(
" {}{:<width$}{} {}{}{}",
dim(),
key,
reset(),
accent(),
value,
reset(),
width = key_width,
)
}
/// Status icon for check results.
///
/// - `pass` → green `✓`
/// - `fail` → red `✗`
/// - `skip` → dim `○`
pub fn status_icon(kind: StatusKind) -> String {
match kind {
StatusKind::Pass => format!("{}\u{2713}{}", success(), reset()),
StatusKind::Fail => format!("{}\u{2717}{}", error(), reset()),
StatusKind::Skip => format!("{}\u{25CB}{}", dim(), reset()),
}
}
/// Kind of status check result.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum StatusKind {
Pass,
Fail,
Skip,
}
/// Top border of a box with an optional label.
///
/// ```text
/// ┌─ label ──────────────────┐
/// ```
pub fn box_top(label: &str, width: usize) -> String {
if label.is_empty() {
let fill = width.saturating_sub(2);
return format!("\u{250C}{}\u{2510}", "\u{2500}".repeat(fill));
}
let label_part = format!(" {} ", label);
// ┌ (1) + ─ (1) + label_part + fill + ┐ (1) = width
let fill = width.saturating_sub(label_part.len() + 3);
format!(
"\u{250C}\u{2500}{}{}{}\u{2510}",
bold(),
label_part,
reset(),
)
.replace("\u{2510}", &format!("{}\u{2510}", "\u{2500}".repeat(fill)))
}
/// Content line inside a box.
///
/// ```text
/// │ content │
/// ```
pub fn box_line(content: &str, width: usize) -> String {
let inner = width.saturating_sub(4); // │ + space + space + │
let padded = if content.len() >= inner {
content.to_string()
} else {
format!("{}{}", content, " ".repeat(inner - content.len()))
};
format!("\u{2502} {} \u{2502}", padded)
}
/// Bottom border of a box.
///
/// ```text
/// └──────────────────────────┘
/// ```
pub fn box_bottom(width: usize) -> String {
let fill = width.saturating_sub(2);
format!("\u{2514}{}\u{2518}", "\u{2500}".repeat(fill))
}
/// Format a check result line for doctor/status commands.
///
/// ```text
/// ✓ Database libsql (connected)
/// ✗ Docker not running — start with: open -a Docker
/// ○ Embeddings disabled
/// ```
pub fn check_line(kind: StatusKind, name: &str, detail: &str, name_width: usize) -> String {
format!(
" {} {:<width$} {}",
status_icon(kind),
name,
detail,
width = name_width,
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn separator_produces_correct_width() {
// In test environment NO_COLOR or non-TTY may be active,
// so strip ANSI to count visible characters.
let s = separator(10);
let visible: String = strip_ansi(&s);
assert_eq!(visible.chars().count(), 10);
}
#[test]
fn kv_line_contains_key_and_value() {
let line = kv_line("model", "gpt-4o", 12);
let visible = strip_ansi(&line);
assert!(visible.contains("model"));
assert!(visible.contains("gpt-4o"));
}
#[test]
fn status_icon_all_kinds() {
// Just verify no panic for each variant
let _ = status_icon(StatusKind::Pass);
let _ = status_icon(StatusKind::Fail);
let _ = status_icon(StatusKind::Skip);
}
#[test]
fn box_drawing() {
let top = box_top("test", 30);
let line = box_line("content", 30);
let bottom = box_bottom(30);
assert!(top.contains('\u{250C}')); // ┌
assert!(line.contains('\u{2502}')); // │
assert!(bottom.contains('\u{2514}')); // └
}
#[test]
fn check_line_formatting() {
let line = check_line(StatusKind::Pass, "Database", "connected", 18);
let visible = strip_ansi(&line);
assert!(visible.contains("Database"));
assert!(visible.contains("connected"));
}
#[test]
fn term_width_in_range() {
let w = term_width();
assert!(w >= 40);
assert!(w <= 120);
}
/// Strip ANSI escape sequences for visible-character counting.
fn strip_ansi(s: &str) -> String {
let mut result = String::new();
let mut in_escape = false;
for c in s.chars() {
if c == '\x1b' {
in_escape = true;
continue;
}
if in_escape {
if c == 'm' {
in_escape = false;
}
continue;
}
result.push(c);
}
result
}
}
+459
View File
@@ -0,0 +1,459 @@
//! Hooks management CLI commands.
//!
//! Lists all discoverable lifecycle hooks from bundled and plugin (WASM
//! capabilities) sources. Plugin discovery uses the same flat-file sidecar
//! layout as the WASM tool/channel loaders (`foo.wasm` + `foo.capabilities.json`).
//!
//! Workspace hooks (`hooks/hooks.json`, `hooks/*.hook.json`) are stored in the
//! database-backed Workspace and require a DB connection to enumerate; this
//! command does not connect to the database, so workspace hooks are omitted.
use std::path::Path;
use clap::Subcommand;
use crate::hooks::bundled::{HookBundleConfig, HookRuleConfig, OutboundWebhookConfig};
use crate::hooks::hook::HookPoint;
const BUNDLED_AUDIT_PRIORITY: u32 = 25;
const DEFAULT_RULE_PRIORITY: u32 = 100;
const DEFAULT_WEBHOOK_PRIORITY: u32 = 300;
#[derive(Subcommand, Debug, Clone)]
pub enum HooksCommand {
/// List discoverable hooks (bundled + plugin; not filtered by active extensions)
List {
/// Show detailed information (hook points, priority, failure mode)
#[arg(short, long)]
verbose: bool,
/// Output as JSON
#[arg(long)]
json: bool,
},
}
/// Run the hooks CLI subcommand.
pub async fn run_hooks_command(
cmd: HooksCommand,
config_path: Option<&Path>,
) -> anyhow::Result<()> {
let config = crate::config::Config::from_env_with_toml(config_path)
.await
.map_err(|e| anyhow::anyhow!("{e:#}"))?;
match cmd {
HooksCommand::List { verbose, json } => cmd_list(&config, verbose, json).await,
}
}
/// Discovered hook information for CLI display.
struct HookInfo {
name: String,
source: String,
kind: String,
points: Vec<HookPoint>,
priority: u32,
failure_mode: String,
}
/// Collect all discoverable hooks from bundled and plugin sources.
async fn discover_hooks(config: &crate::config::Config) -> Vec<HookInfo> {
let mut hooks = Vec::new();
// 1. Bundled hooks (hardcoded)
hooks.push(HookInfo {
name: "builtin.audit_log".to_string(),
source: "bundled".to_string(),
kind: "audit".to_string(),
points: vec![
HookPoint::BeforeInbound,
HookPoint::BeforeToolCall,
HookPoint::BeforeOutbound,
HookPoint::OnSessionStart,
HookPoint::OnSessionEnd,
HookPoint::TransformResponse,
],
priority: BUNDLED_AUDIT_PRIORITY,
failure_mode: "fail_open".to_string(),
});
// 2. Plugin hooks from WASM capabilities sidecar files
let wasm_tools_dir = &config.wasm.tools_dir;
let wasm_channels_dir = &config.channels.wasm_channels_dir;
collect_plugin_hooks(&mut hooks, wasm_tools_dir, "tool").await;
collect_plugin_hooks(&mut hooks, wasm_channels_dir, "channel").await;
// Note: workspace hooks (hooks/hooks.json, hooks/*.hook.json) are stored
// in the database-backed Workspace and require a DB connection to list.
// Sort by priority then name for stable output
hooks.sort_by(|a, b| a.priority.cmp(&b.priority).then(a.name.cmp(&b.name)));
hooks
}
/// Scan a WASM directory for `*.capabilities.json` sidecar files containing hook
/// definitions.
///
/// Uses the same flat-file layout as the real WASM loaders:
/// ```text
/// ~/.ironclaw/tools/
/// ├── slack.wasm
/// ├── slack.capabilities.json <- hooks section parsed here
/// ├── github.wasm
/// └── github.capabilities.json
/// ```
async fn collect_plugin_hooks(hooks: &mut Vec<HookInfo>, dir: &Path, plugin_type: &str) {
if !dir.exists() {
return;
}
let mut entries = match tokio::fs::read_dir(dir).await {
Ok(entries) => entries,
Err(_) => return,
};
while let Ok(Some(entry)) = entries.next_entry().await {
let path = entry.path();
// Match only *.capabilities.json sidecar files (flat layout)
let file_name = match path.file_name().and_then(|n| n.to_str()) {
Some(n) => n.to_string(),
None => continue,
};
if !file_name.ends_with(".capabilities.json") {
continue;
}
// Extract tool/channel name: "slack.capabilities.json" -> "slack"
let name = match file_name.strip_suffix(".capabilities.json") {
Some(n) if !n.is_empty() => n.to_string(),
_ => continue,
};
let bytes = match tokio::fs::read(&path).await {
Ok(b) => b,
Err(_) => continue,
};
let value: serde_json::Value = match serde_json::from_slice(&bytes) {
Ok(v) => v,
Err(_) => continue,
};
// Match the same extraction logic as bootstrap: check "hooks" key
// at root or nested under "capabilities.hooks".
let hooks_section = value
.get("hooks")
.or_else(|| value.get("capabilities").and_then(|c| c.get("hooks")));
let Some(hooks_value) = hooks_section else {
continue;
};
let bundle = match HookBundleConfig::from_value(hooks_value) {
Ok(b) => b,
Err(_) => continue,
};
let source = format!("plugin.{plugin_type}:{name}");
for rule in &bundle.rules {
hooks.push(hook_info_from_rule(&source, rule));
}
for webhook in &bundle.outbound_webhooks {
hooks.push(hook_info_from_webhook(&source, webhook));
}
}
}
fn hook_info_from_rule(source: &str, rule: &HookRuleConfig) -> HookInfo {
let scoped_name = format!("{source}::{}", rule.name);
HookInfo {
name: scoped_name,
source: source.to_string(),
kind: if rule.reject_reason.is_some() {
"reject".to_string()
} else {
"rule".to_string()
},
points: rule.points.clone(),
priority: rule.priority.unwrap_or(DEFAULT_RULE_PRIORITY),
failure_mode: rule
.failure_mode
.as_ref()
.map(|m| format!("{m:?}"))
.unwrap_or_else(|| "fail_open".to_string()),
}
}
fn hook_info_from_webhook(source: &str, webhook: &OutboundWebhookConfig) -> HookInfo {
let scoped_name = format!("{source}::{}", webhook.name);
HookInfo {
name: scoped_name,
source: source.to_string(),
kind: "webhook".to_string(),
points: webhook.points.clone(),
priority: webhook.priority.unwrap_or(DEFAULT_WEBHOOK_PRIORITY),
failure_mode: "fail_open".to_string(),
}
}
/// List all discovered hooks.
async fn cmd_list(config: &crate::config::Config, verbose: bool, json: bool) -> anyhow::Result<()> {
let hooks = discover_hooks(config).await;
if json {
let entries: Vec<serde_json::Value> = hooks
.iter()
.map(|h| {
let mut v = serde_json::json!({
"name": h.name,
"source": h.source,
"kind": h.kind,
"priority": h.priority,
"points": h.points.iter().map(|p| p.as_str()).collect::<Vec<_>>(),
});
if verbose {
v["failure_mode"] = serde_json::json!(h.failure_mode);
}
v
})
.collect();
println!(
"{}",
serde_json::to_string_pretty(&entries).unwrap_or_else(|_| "[]".to_string())
);
return Ok(());
}
if hooks.is_empty() {
println!("No hooks found.");
return Ok(());
}
println!("Discovered {} hook(s):\n", hooks.len());
for h in &hooks {
if verbose {
let points_str: Vec<&str> = h.points.iter().map(|p| p.as_str()).collect();
println!(" {}", h.name);
println!(" Source: {}", h.source);
println!(" Kind: {}", h.kind);
println!(" Priority: {}", h.priority);
println!(" Points: {}", points_str.join(", "));
println!(" Failure mode: {}", h.failure_mode);
println!();
} else {
let points_str: Vec<&str> = h.points.iter().map(|p| p.as_str()).collect();
println!(
" {:<40} [{:<7}] pri={:<3} {}",
h.name,
h.kind,
h.priority,
points_str.join(", ")
);
}
}
if !verbose {
println!();
println!(
"Use --verbose for details. Workspace hooks (DB-stored) are not listed without a database connection."
);
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Write;
#[test]
fn hook_info_from_rule_basic() {
let rule = HookRuleConfig {
name: "test-rule".to_string(),
points: vec![HookPoint::BeforeInbound],
priority: Some(50),
failure_mode: None,
timeout_ms: None,
when_regex: None,
reject_reason: None,
replacements: vec![],
prepend: None,
append: None,
};
let info = hook_info_from_rule("plugin.tool:my_tool", &rule);
assert_eq!(info.name, "plugin.tool:my_tool::test-rule");
assert_eq!(info.source, "plugin.tool:my_tool");
assert_eq!(info.kind, "rule");
assert_eq!(info.priority, 50);
}
#[test]
fn hook_info_from_rule_reject() {
let rule = HookRuleConfig {
name: "blocker".to_string(),
points: vec![HookPoint::BeforeInbound, HookPoint::BeforeToolCall],
priority: None,
failure_mode: None,
timeout_ms: None,
when_regex: Some("bad_pattern".to_string()),
reject_reason: Some("blocked".to_string()),
replacements: vec![],
prepend: None,
append: None,
};
let info = hook_info_from_rule("workspace:hooks/block.hook.json", &rule);
assert_eq!(info.kind, "reject");
assert_eq!(info.priority, DEFAULT_RULE_PRIORITY);
}
#[test]
fn hook_info_from_webhook_basic() {
let webhook = OutboundWebhookConfig {
name: "notify".to_string(),
points: vec![HookPoint::BeforeOutbound],
url: "https://example.com/hook".to_string(),
headers: Default::default(),
timeout_ms: None,
priority: Some(200),
max_in_flight: None,
};
let info = hook_info_from_webhook("plugin.tool:logger", &webhook);
assert_eq!(info.name, "plugin.tool:logger::notify");
assert_eq!(info.kind, "webhook");
assert_eq!(info.priority, 200);
}
#[tokio::test]
async fn discover_plugin_hooks_flat_layout() {
let dir = tempfile::tempdir().expect("create temp dir");
// Create a sidecar capabilities file with hooks (flat layout)
let caps = serde_json::json!({
"hooks": {
"rules": [
{
"name": "redact-keys",
"points": ["beforeOutbound"],
"replacements": [
{"pattern": "sk-[a-zA-Z0-9]+", "replacement": "[REDACTED]"}
]
}
],
"outbound_webhooks": [
{
"name": "log-events",
"points": ["beforeInbound"],
"url": "https://example.com/events"
}
]
}
});
let mut f =
std::fs::File::create(dir.path().join("slack.capabilities.json")).expect("create file");
f.write_all(serde_json::to_string(&caps).unwrap().as_bytes())
.expect("write");
// Also create a .wasm file (not required for discovery, but realistic)
std::fs::File::create(dir.path().join("slack.wasm")).expect("create wasm");
// A capabilities file without hooks should be skipped
let no_hooks = serde_json::json!({"http": {"allowlist": []}});
let mut f2 = std::fs::File::create(dir.path().join("github.capabilities.json"))
.expect("create file");
f2.write_all(serde_json::to_string(&no_hooks).unwrap().as_bytes())
.expect("write");
let mut hooks = Vec::new();
collect_plugin_hooks(&mut hooks, dir.path(), "tool").await;
assert_eq!(hooks.len(), 2, "should find 1 rule + 1 webhook");
assert_eq!(hooks[0].name, "plugin.tool:slack::redact-keys");
assert_eq!(hooks[0].kind, "rule");
assert_eq!(hooks[1].name, "plugin.tool:slack::log-events");
assert_eq!(hooks[1].kind, "webhook");
}
#[tokio::test]
async fn discover_plugin_hooks_nested_capabilities() {
let dir = tempfile::tempdir().expect("create temp dir");
// Channel-style capabilities with hooks nested under "capabilities"
let caps = serde_json::json!({
"type": "channel",
"capabilities": {
"hooks": {
"rules": [
{
"name": "filter-spam",
"points": ["beforeInbound"],
"when_regex": "buy now",
"reject_reason": "spam detected"
}
]
}
}
});
let mut f = std::fs::File::create(dir.path().join("telegram.capabilities.json"))
.expect("create file");
f.write_all(serde_json::to_string(&caps).unwrap().as_bytes())
.expect("write");
let mut hooks = Vec::new();
collect_plugin_hooks(&mut hooks, dir.path(), "channel").await;
assert_eq!(hooks.len(), 1);
assert_eq!(hooks[0].name, "plugin.channel:telegram::filter-spam");
assert_eq!(hooks[0].kind, "reject");
assert_eq!(hooks[0].source, "plugin.channel:telegram");
}
#[tokio::test]
async fn discover_plugin_hooks_empty_dir() {
let dir = tempfile::tempdir().expect("create temp dir");
let mut hooks = Vec::new();
collect_plugin_hooks(&mut hooks, dir.path(), "tool").await;
assert!(hooks.is_empty());
}
#[tokio::test]
async fn discover_plugin_hooks_nonexistent_dir() {
let mut hooks = Vec::new();
collect_plugin_hooks(&mut hooks, Path::new("/nonexistent/path"), "tool").await;
assert!(hooks.is_empty());
}
#[tokio::test]
async fn discover_plugin_hooks_skips_subdirectories() {
let dir = tempfile::tempdir().expect("create temp dir");
// Create a subdirectory with capabilities.json inside (old broken layout)
// This should NOT be discovered — only flat sidecar files are valid.
let sub = dir.path().join("my_tool");
std::fs::create_dir_all(&sub).expect("create subdir");
let caps =
serde_json::json!({"hooks": {"rules": [{"name": "x", "points": ["beforeInbound"]}]}});
let mut f = std::fs::File::create(sub.join("capabilities.json")).expect("create file");
f.write_all(serde_json::to_string(&caps).unwrap().as_bytes())
.expect("write");
let mut hooks = Vec::new();
collect_plugin_hooks(&mut hooks, dir.path(), "tool").await;
// The subdirectory layout should be ignored
assert!(
hooks.is_empty(),
"subdirectory capabilities.json should not be discovered"
);
}
}
+18 -3
View File
@@ -18,6 +18,8 @@ mod channels;
mod completion;
mod config;
mod doctor;
pub mod fmt;
mod hooks;
#[cfg(feature = "import")]
pub mod import;
mod logs;
@@ -36,6 +38,7 @@ pub use channels::{ChannelsCommand, run_channels_command};
pub use completion::Completion;
pub use config::{ConfigCommand, run_config_command};
pub use doctor::run_doctor_command;
pub use hooks::{HooksCommand, run_hooks_command};
#[cfg(feature = "import")]
pub use import::{ImportCommand, run_import_command};
pub use logs::{LogsCommand, run_logs_command};
@@ -109,16 +112,20 @@ pub enum Command {
skip_auth: bool,
/// Reconfigure channels only
#[arg(long, conflicts_with_all = ["provider_only", "quick"])]
#[arg(long, conflicts_with_all = ["provider_only", "quick", "step"], help = "Deprecated: use --step channels")]
channels_only: bool,
/// Reconfigure LLM provider and model only
#[arg(long, conflicts_with_all = ["channels_only", "quick"])]
#[arg(long, conflicts_with_all = ["channels_only", "quick", "step"], help = "Deprecated: use --step provider")]
provider_only: bool,
/// Quick setup: auto-defaults everything except LLM provider and model
#[arg(long, conflicts_with_all = ["channels_only", "provider_only"])]
#[arg(long, conflicts_with_all = ["channels_only", "provider_only", "step"])]
quick: bool,
/// Run only specific setup steps (comma-separated: provider, channels, model, database, security)
#[arg(long, value_delimiter = ',', conflicts_with_all = ["channels_only", "provider_only", "quick"])]
step: Vec<String>,
},
/// Manage configuration settings
@@ -202,6 +209,14 @@ pub enum Command {
)]
Skills(SkillsCommand),
/// Manage lifecycle hooks
#[command(
subcommand,
about = "Manage lifecycle hooks",
long_about = "List and inspect lifecycle hooks (bundled, plugin, workspace).\nExamples:\n ironclaw hooks list\n ironclaw hooks list --verbose\n ironclaw hooks list --json"
)]
Hooks(HooksCommand),
/// Probe external dependencies and validate configuration
#[command(
about = "Run diagnostics",
+83 -18
View File
@@ -579,23 +579,27 @@ pub fn encode_hosted_oauth_state(flow_id: &str, instance_name: Option<&str>) ->
/// Decode hosted OAuth state in either the new versioned format or the
/// legacy `instance:nonce`/`nonce` forms.
pub fn decode_hosted_oauth_state(state: &str) -> Result<DecodedHostedOAuthState, String> {
if let Some(rest) = state.strip_prefix(&format!("{HOSTED_STATE_PREFIX}."))
&& let Some((payload_b64, checksum)) = rest.rsplit_once('.')
&& let Ok(payload_json) = URL_SAFE_NO_PAD.decode(payload_b64)
{
if let Some(rest) = state.strip_prefix(&format!("{HOSTED_STATE_PREFIX}.")) {
let (payload_b64, checksum) = rest
.rsplit_once('.')
.ok_or("Hosted OAuth versioned state missing checksum separator")?;
let payload_json = URL_SAFE_NO_PAD
.decode(payload_b64)
.map_err(|e| format!("Hosted OAuth versioned state base64 decode failed: {e}"))?;
let expected_checksum = hosted_state_checksum(&payload_json);
if checksum != expected_checksum {
return Err("Hosted OAuth state checksum mismatch".to_string());
}
if let Ok(payload) = serde_json::from_slice::<HostedOAuthStatePayload>(&payload_json)
&& !payload.flow_id.trim().is_empty()
{
return Ok(DecodedHostedOAuthState {
flow_id: payload.flow_id,
instance_name: payload.instance_name.filter(|v| !v.is_empty()),
is_legacy: false,
});
let payload: HostedOAuthStatePayload = serde_json::from_slice(&payload_json)
.map_err(|e| format!("Hosted OAuth versioned state JSON parse failed: {e}"))?;
if payload.flow_id.trim().is_empty() {
return Err("Hosted OAuth versioned state has empty flow_id".to_string());
}
return Ok(DecodedHostedOAuthState {
flow_id: payload.flow_id,
instance_name: payload.instance_name.filter(|v| !v.is_empty()),
is_legacy: false,
});
}
if let Some((instance_name, flow_id)) = state.split_once(':') {
@@ -1187,14 +1191,14 @@ mod tests {
}
#[test]
fn test_decode_hosted_oauth_state_falls_back_for_non_envelope_ic2_prefix() {
fn test_decode_hosted_oauth_state_rejects_non_envelope_ic2_prefix() {
use crate::cli::oauth_defaults::decode_hosted_oauth_state;
let decoded =
decode_hosted_oauth_state("ic2.provider-owned-state").expect("prefixed fallback");
assert_eq!(decoded.flow_id, "ic2.provider-owned-state");
assert_eq!(decoded.instance_name, None);
assert!(decoded.is_legacy);
// "ic2." prefix must parse as a valid versioned envelope — never fall
// through to legacy handling, which would use the full malformed
// envelope as the flow_id and break OAuth callback lookup (#1441).
decode_hosted_oauth_state("ic2.provider-owned-state")
.expect_err("ic2-prefixed non-envelope state should fail");
}
#[test]
@@ -1244,4 +1248,65 @@ mod tests {
assert!(result.url.contains("code_challenge="));
assert!(result.code_verifier.is_some());
}
/// Malformed `ic2.*` states must return Err, never fall through to legacy
/// handling where the full envelope would be used as the flow_id (#1441).
#[test]
fn test_decode_versioned_state_rejects_malformed_envelopes() {
use crate::cli::oauth_defaults::decode_hosted_oauth_state;
// Missing checksum separator (no second dot after prefix)
let err =
decode_hosted_oauth_state("ic2.nodots").expect_err("missing separator should fail");
assert!(
err.contains("checksum separator"),
"unexpected error: {err}"
);
// Bad base64 payload
let err = decode_hosted_oauth_state("ic2.!!!badbase64!!!.fakechecksum")
.expect_err("bad base64 should fail");
assert!(err.contains("base64"), "unexpected error: {err}");
// Valid base64 but not JSON: use correct checksum so we exercise JSON parsing
use base64::Engine;
use sha2::Digest;
let not_json_bytes = b"not json";
let not_json_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(not_json_bytes);
let digest = sha2::Sha256::digest(not_json_bytes);
let checksum = base64::engine::general_purpose::URL_SAFE_NO_PAD
.encode(&digest[..super::HOSTED_STATE_CHECKSUM_BYTES]);
let err = decode_hosted_oauth_state(&format!("ic2.{not_json_b64}.{checksum}"))
.expect_err("non-JSON payload should fail with JSON parse error");
assert!(
err.contains("JSON"),
"unexpected error (expected JSON parse failure): {err}"
);
}
/// Round-trip: encode_hosted_oauth_state(nonce) → decode → flow_id == nonce.
/// Ensures the registration key and lookup key are always identical (#1441).
#[test]
fn test_oauth_flow_key_round_trip_consistency() {
use crate::cli::oauth_defaults::{decode_hosted_oauth_state, encode_hosted_oauth_state};
let nonce = "test-nonce-abc123";
let encoded = encode_hosted_oauth_state(nonce, Some("my-instance"));
let decoded = decode_hosted_oauth_state(&encoded).expect("round-trip decode");
assert_eq!(
decoded.flow_id, nonce,
"flow_id must match the original nonce"
);
assert_eq!(decoded.instance_name.as_deref(), Some("my-instance"));
assert!(!decoded.is_legacy);
// Also test without instance name
let encoded_no_instance = encode_hosted_oauth_state(nonce, None);
let decoded_no_instance =
decode_hosted_oauth_state(&encoded_no_instance).expect("round-trip without instance");
assert_eq!(decoded_no_instance.flow_id, nonce);
assert_eq!(decoded_no_instance.instance_name, None);
assert!(!decoded_no_instance.is_legacy);
}
}
@@ -19,6 +19,7 @@ Commands:
pairing Manage DM pairing
service Manage OS service
skills Manage skills
hooks Manage lifecycle hooks
doctor Run diagnostics
logs View and manage gateway logs
status Show system status
@@ -19,6 +19,7 @@ Commands:
pairing Manage DM pairing
service Manage OS service
skills Manage skills
hooks Manage lifecycle hooks
doctor Run diagnostics
logs View and manage gateway logs
status Show system status
@@ -22,6 +22,7 @@ Commands:
pairing Manage DM pairing
service Manage OS service
skills Manage skills
hooks Manage lifecycle hooks
doctor Run diagnostics
logs View and manage gateway logs
status Show system status
@@ -22,6 +22,7 @@ Commands:
pairing Manage DM pairing
service Manage OS service
skills Manage skills
hooks Manage lifecycle hooks
doctor Run diagnostics
logs View and manage gateway logs
status Show system status
+57 -48
View File
@@ -6,6 +6,7 @@
use std::path::PathBuf;
use crate::bootstrap::ironclaw_base_dir;
use crate::cli::fmt;
use crate::settings::Settings;
/// Load settings from JSON and TOML config files, matching the runtime
@@ -38,22 +39,25 @@ fn load_settings_from(json_path: &std::path::Path, toml_path: &std::path::Path)
pub async fn run_status_command() -> anyhow::Result<()> {
let settings = load_settings();
println!("IronClaw Status");
println!("===============\n");
println!();
println!(" {}IronClaw Status{}", fmt::bold(), fmt::reset());
println!();
// Version
println!(
" Version: {} v{}",
env!("CARGO_PKG_NAME"),
env!("CARGO_PKG_VERSION")
"{}",
fmt::kv_line(
"Version",
&format!("{} v{}", env!("CARGO_PKG_NAME"), env!("CARGO_PKG_VERSION")),
12,
)
);
// Database
print!(" Database: ");
let db_backend = std::env::var("DATABASE_BACKEND")
.ok()
.unwrap_or_else(|| "postgres".to_string());
match db_backend.as_str() {
let db_value = match db_backend.as_str() {
"libsql" | "turso" | "sqlite" => {
let path = std::env::var("LIBSQL_PATH")
.map(std::path::PathBuf::from)
@@ -64,77 +68,77 @@ pub async fn run_status_command() -> anyhow::Result<()> {
} else {
""
};
println!("libSQL ({}{})", path.display(), turso);
format!("libSQL ({}{})", path.display(), turso)
} else {
println!("libSQL (file missing: {})", path.display());
format!("libSQL (file missing: {})", path.display())
}
}
_ => {
if std::env::var("DATABASE_URL").is_ok() {
match check_database().await {
Ok(()) => println!("connected (PostgreSQL)"),
Err(e) => println!("error ({})", e),
Ok(()) => "connected (PostgreSQL)".to_string(),
Err(e) => format!("error ({})", e),
}
} else {
println!("not configured");
"not configured".to_string()
}
}
}
};
println!("{}", fmt::kv_line("Database", &db_value, 12));
// Session / Auth
print!(" Session: ");
let session_path = crate::config::llm::default_session_path();
if session_path.exists() {
println!("found ({})", session_path.display());
let session_value = if session_path.exists() {
format!("found ({})", session_path.display())
} else {
println!("not found (run `ironclaw onboard`)");
}
"not found (run `ironclaw onboard`)".to_string()
};
println!("{}", fmt::kv_line("Session", &session_value, 12));
// Secrets (auto-detect from env only; skip keychain probe to avoid
// triggering macOS system password dialogs on a simple status check)
print!(" Secrets: ");
if std::env::var("SECRETS_MASTER_KEY").is_ok() {
println!("configured (env)");
let secrets_value = if std::env::var("SECRETS_MASTER_KEY").is_ok() {
"configured (env)".to_string()
} else {
// We don't probe the keychain here because get_generic_password()
// triggers macOS unlock+authorization dialogs, which is bad UX for
// a read-only status command. If onboarding completed with keychain
// storage, the key is there; we just can't cheaply verify it.
println!("env not set (keychain may be configured)");
}
"env not set (keychain may be configured)".to_string()
};
println!("{}", fmt::kv_line("Secrets", &secrets_value, 12));
// Embeddings
print!(" Embeddings: ");
let emb_enabled = settings.embeddings.enabled
|| std::env::var("OPENAI_API_KEY").is_ok()
|| std::env::var("EMBEDDING_ENABLED")
.map(|v| v == "true")
.unwrap_or(false);
if emb_enabled {
println!(
let emb_value = if emb_enabled {
format!(
"enabled (provider: {}, model: {})",
settings.embeddings.provider, settings.embeddings.model
);
)
} else {
println!("disabled");
}
"disabled".to_string()
};
println!("{}", fmt::kv_line("Embeddings", &emb_value, 12));
// WASM tools
print!(" WASM Tools: ");
let tools_dir = settings
.wasm
.tools_dir
.clone()
.unwrap_or_else(default_tools_dir);
if tools_dir.exists() {
let tools_value = if tools_dir.exists() {
let count = count_wasm_files(&tools_dir);
println!("{} installed ({})", count, tools_dir.display());
format!("{} installed ({})", count, tools_dir.display())
} else {
println!("directory not found ({})", tools_dir.display());
}
format!("directory not found ({})", tools_dir.display())
};
println!("{}", fmt::kv_line("WASM Tools", &tools_value, 12));
// WASM channels
print!(" Channels: ");
let channels_dir = settings
.channels
.wasm_channels_dir
@@ -153,35 +157,40 @@ pub async fn run_status_command() -> anyhow::Result<()> {
channel_info.push(format!("{} wasm", wasm_count));
}
}
println!("{}", channel_info.join(", "));
println!("{}", fmt::kv_line("Channels", &channel_info.join(", "), 12));
// Heartbeat
print!(" Heartbeat: ");
let hb_enabled = settings.heartbeat.enabled
|| std::env::var("HEARTBEAT_ENABLED")
.map(|v| v == "true")
.unwrap_or(false);
if hb_enabled {
println!("enabled (interval: {}s)", settings.heartbeat.interval_secs);
let hb_value = if hb_enabled {
format!("enabled (interval: {}s)", settings.heartbeat.interval_secs)
} else {
println!("disabled");
}
"disabled".to_string()
};
println!("{}", fmt::kv_line("Heartbeat", &hb_value, 12));
// MCP servers
print!(" MCP Servers: ");
match crate::tools::mcp::config::load_mcp_servers().await {
let mcp_value = match crate::tools::mcp::config::load_mcp_servers().await {
Ok(servers) => {
let enabled = servers.servers.iter().filter(|s| s.enabled).count();
let total = servers.servers.len();
println!("{} enabled / {} configured", enabled, total);
format!("{} enabled / {} configured", enabled, total)
}
Err(_) => println!("none configured"),
}
Err(_) => "none configured".to_string(),
};
println!("{}", fmt::kv_line("MCP Servers", &mcp_value, 12));
// Config path
println!();
println!(
"\n Config: {}",
crate::bootstrap::ironclaw_env_path().display()
"{}",
fmt::kv_line(
"Config",
&crate::bootstrap::ironclaw_env_path().display().to_string(),
12,
)
);
Ok(())
+125 -3
View File
@@ -9,6 +9,7 @@ use crate::llm::config::*;
use crate::llm::registry::{ProviderProtocol, ProviderRegistry};
use crate::llm::session::SessionConfig;
use crate::settings::Settings;
impl LlmConfig {
/// Create a test-friendly config without reading env vars.
#[cfg(feature = "libsql")]
@@ -37,6 +38,7 @@ impl LlmConfig {
},
provider: None,
bedrock: None,
gemini_oauth: None,
openai_codex: None,
request_timeout_secs: 120,
cheap_model: None,
@@ -73,11 +75,16 @@ impl LlmConfig {
backend_lower == "nearai" || backend_lower == "near_ai" || backend_lower == "near";
let is_bedrock =
backend_lower == "bedrock" || backend_lower == "aws_bedrock" || backend_lower == "aws";
let is_gemini_oauth = backend_lower == "gemini_oauth" || backend_lower == "gemini-oauth";
let is_openai_codex = backend_lower == "openai_codex"
|| backend_lower == "openai-codex"
|| backend_lower == "codex";
if !is_nearai && !is_bedrock && !is_openai_codex && registry.find(&backend_lower).is_none()
if !is_nearai
&& !is_bedrock
&& !is_gemini_oauth
&& !is_openai_codex
&& registry.find(&backend_lower).is_none()
{
tracing::warn!(
"Unknown LLM backend '{}'. Will attempt as openai_compatible fallback.",
@@ -131,8 +138,8 @@ impl LlmConfig {
smart_routing_cascade: parse_optional_env("SMART_ROUTING_CASCADE", true)?,
};
// Resolve registry provider config (for non-NearAI, non-Bedrock, non-Codex backends)
let provider = if is_nearai || is_bedrock || is_openai_codex {
// Resolve registry provider config (for non-NearAI, non-Bedrock, non-Gemini, non-Codex backends)
let provider = if is_nearai || is_bedrock || is_gemini_oauth || is_openai_codex {
None
} else {
Some(Self::resolve_registry_provider(
@@ -213,6 +220,19 @@ impl LlmConfig {
let request_timeout_secs = parse_optional_env("LLM_REQUEST_TIMEOUT_SECS", 120)?;
let gemini_oauth = if backend_lower == "gemini_oauth" || backend_lower == "gemini-oauth" {
let model = Self::resolve_model("GEMINI_MODEL", settings, "gemini-2.5-flash")?;
let credentials_path = optional_env("GEMINI_CREDENTIALS_PATH")?
.map(PathBuf::from)
.unwrap_or_else(GeminiOauthConfig::default_credentials_path);
Some(GeminiOauthConfig {
model,
credentials_path,
})
} else {
None
};
// Generic cheap model (works with any backend).
// Falls back to NearAI-specific cheap_model in provider chain logic.
let cheap_model = optional_env("LLM_CHEAP_MODEL")?;
@@ -226,6 +246,8 @@ impl LlmConfig {
"nearai".to_string()
} else if is_bedrock {
"bedrock".to_string()
} else if is_gemini_oauth {
"gemini_oauth".to_string()
} else if is_openai_codex {
"openai_codex".to_string()
} else if let Some(ref p) = provider {
@@ -237,6 +259,7 @@ impl LlmConfig {
nearai,
provider,
bedrock,
gemini_oauth,
openai_codex,
request_timeout_secs,
cheap_model,
@@ -389,6 +412,14 @@ impl LlmConfig {
} else {
Vec::new()
};
let extra_headers = if canonical_id == "github_copilot" {
merge_extra_headers(
crate::llm::github_copilot_auth::default_headers(),
extra_headers,
)
} else {
extra_headers
};
// Resolve OAuth token (Anthropic-specific: `claude login` flow).
// Only check for OAuth token when the provider is actually Anthropic.
@@ -473,6 +504,26 @@ fn parse_extra_headers(val: &str) -> Result<Vec<(String, String)>, ConfigError>
Ok(headers)
}
fn merge_extra_headers(
defaults: Vec<(String, String)>,
overrides: Vec<(String, String)>,
) -> Vec<(String, String)> {
let mut merged = Vec::new();
let mut positions = std::collections::HashMap::<String, usize>::new();
for (key, value) in defaults.into_iter().chain(overrides) {
let normalized = key.to_ascii_lowercase();
if let Some(existing_index) = positions.get(&normalized).copied() {
merged[existing_index] = (key, value);
} else {
positions.insert(normalized, merged.len());
merged.push((key, value));
}
}
merged
}
/// Get the default session file path (~/.ironclaw/session.json).
pub fn default_session_path() -> PathBuf {
ironclaw_base_dir().join("session.json")
@@ -604,6 +655,29 @@ mod tests {
);
}
#[test]
fn merge_extra_headers_prefers_overrides_case_insensitively() {
let merged = merge_extra_headers(
vec![
("User-Agent".to_string(), "default-agent".to_string()),
("X-Test".to_string(), "default".to_string()),
],
vec![
("user-agent".to_string(), "override-agent".to_string()),
("X-Extra".to_string(), "present".to_string()),
],
);
assert_eq!(
merged,
vec![
("user-agent".to_string(), "override-agent".to_string()),
("X-Test".to_string(), "default".to_string()),
("X-Extra".to_string(), "present".to_string()),
]
);
}
/// Clear all ollama-related env vars.
fn clear_ollama_env() {
// SAFETY: Only called under ENV_MUTEX in tests.
@@ -756,6 +830,54 @@ mod tests {
assert_eq!(provider.protocol, ProviderProtocol::OpenAiCompletions);
}
#[test]
fn registry_provider_resolves_github_copilot_alias() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::set_var("LLM_BACKEND", "github-copilot");
std::env::set_var("GITHUB_COPILOT_TOKEN", "gho_test_token");
std::env::set_var(
"GITHUB_COPILOT_EXTRA_HEADERS",
"Copilot-Integration-Id:custom-chat,X-Test:enabled",
);
}
let settings = Settings::default();
let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed");
assert_eq!(cfg.backend, "github_copilot");
let provider = cfg.provider.expect("provider config should be present");
assert_eq!(provider.provider_id, "github_copilot");
assert_eq!(provider.base_url, "https://api.githubcopilot.com");
assert_eq!(provider.model, "gpt-4o");
assert!(
provider
.extra_headers
.iter()
.any(|(key, value)| { key == "Copilot-Integration-Id" && value == "custom-chat" })
);
assert!(
provider
.extra_headers
.iter()
.any(|(key, value)| key == "User-Agent" && value == "GitHubCopilotChat/0.26.7")
);
assert!(
provider
.extra_headers
.iter()
.any(|(key, value)| key == "X-Test" && value == "enabled")
);
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::remove_var("LLM_BACKEND");
std::env::remove_var("GITHUB_COPILOT_TOKEN");
std::env::remove_var("GITHUB_COPILOT_EXTRA_HEADERS");
}
}
#[test]
fn nearai_backend_has_no_registry_provider() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
+2 -2
View File
@@ -56,8 +56,8 @@ pub use self::tunnel::TunnelConfig;
pub use self::wasm::WasmConfig;
pub use self::workspace::WorkspaceConfig;
pub use crate::llm::config::{
BedrockConfig, CacheRetention, LlmConfig, NearAiConfig, OAUTH_PLACEHOLDER, OpenAiCodexConfig,
RegistryProviderConfig,
BedrockConfig, CacheRetention, GeminiOauthConfig, LlmConfig, NearAiConfig, OAUTH_PLACEHOLDER,
OpenAiCodexConfig, RegistryProviderConfig,
};
pub use crate::llm::session::SessionConfig;
+9 -6
View File
@@ -89,7 +89,9 @@ impl TranscriptionConfig {
}
/// Create the transcription provider if enabled and configured.
pub fn create_provider(&self) -> Option<Box<dyn crate::transcription::TranscriptionProvider>> {
pub fn create_provider(
&self,
) -> Option<Box<dyn crate::llm::transcription::TranscriptionProvider>> {
if !self.enabled {
return None;
}
@@ -103,10 +105,11 @@ impl TranscriptionConfig {
"Audio transcription enabled via Chat Completions API"
);
let mut provider = crate::transcription::ChatCompletionsTranscriptionProvider::new(
api_key.clone(),
)
.with_model(&self.model);
let mut provider =
crate::llm::transcription::ChatCompletionsTranscriptionProvider::new(
api_key.clone(),
)
.with_model(&self.model);
if let Some(ref base_url) = self.base_url {
provider = provider.with_base_url(base_url);
@@ -121,7 +124,7 @@ impl TranscriptionConfig {
);
let mut provider =
crate::transcription::OpenAiWhisperProvider::new(api_key.clone())
crate::llm::transcription::OpenAiWhisperProvider::new(api_key.clone())
.with_model(&self.model);
if let Some(ref base_url) = self.base_url {
+1 -1
View File
@@ -36,7 +36,7 @@ pub(crate) fn resolve_embedding_dimension() -> Option<usize> {
.unwrap_or(false);
if !enabled {
tracing::info!("Vector index setup skipped (EMBEDDING_ENABLED not set in env)");
tracing::debug!("Vector index setup skipped (EMBEDDING_ENABLED not set in env)");
return None;
}
+1 -1
View File
@@ -97,7 +97,7 @@ pub async fn connect_with_handles(
.map_err(|e| DatabaseError::Pool(e.to_string()))?
};
backend.run_migrations().await?;
tracing::info!("libSQL database connected and migrations applied");
tracing::debug!("libSQL database connected and migrations applied");
handles.libsql_db = Some(backend.shared_db());
+508 -55
View File
@@ -107,6 +107,21 @@ struct ChannelRuntimeState {
wasm_channel_owner_ids: std::collections::HashMap<String, i64>,
}
/// Setup schema returned to web UI for extension configuration.
pub struct ExtensionSetupSchema {
pub secrets: Vec<crate::channels::web::types::SecretFieldInfo>,
pub fields: Vec<crate::channels::web::types::SetupFieldInfo>,
}
/// Only these global (non-namespaced) setting paths may be written by extension
/// setup fields. Everything else must be under `extensions.<name>.*`.
const ALLOWED_GLOBAL_SETUP_SETTING_PATHS: &[&str] = &[
"llm_backend",
"selected_model",
"ollama_base_url",
"openai_compatible_base_url",
];
#[cfg(test)]
type TestWasmChannelLoader =
Arc<dyn Fn(&str) -> Result<LoadedChannel, ExtensionError> + Send + Sync>;
@@ -937,6 +952,31 @@ impl ExtensionManager {
&self.secrets
}
/// Inject a pre-created MCP client (from startup loading) into the manager.
///
/// Startup-loaded MCP clients register their tools in `ToolRegistry` but are
/// otherwise dropped. This method stores the client so that `list()` reports
/// accurate "connected" status and reconnection/session management works.
pub(crate) async fn inject_mcp_client(
&self,
name: String,
client: Arc<crate::tools::mcp::McpClient>,
) {
if name.is_empty() {
tracing::warn!("inject_mcp_client called with empty name; ignoring");
return;
}
if let Err(e) = Self::validate_extension_name(&name) {
tracing::warn!(
error = %e,
name = %name,
"inject_mcp_client called with invalid name; ignoring"
);
return;
}
self.mcp_clients.write().await.insert(name, client);
}
/// Register channel names that were loaded at startup.
/// Called after WASM channels are loaded so `list()` reports accurate active status.
pub async fn set_active_channels(&self, names: Vec<String>) {
@@ -3316,6 +3356,46 @@ impl ExtensionManager {
return ToolAuthState::NoAuth;
};
let saved_fields = self.load_tool_setup_fields(name).await.unwrap_or_default();
let setup_is_complete = if let Some(setup) = &cap_file.setup {
let secrets_ready = futures::future::join_all(
setup
.required_secrets
.iter()
.filter(|s| !s.optional)
.filter(|s| !Self::is_auto_resolved_oauth_field(&s.name, &cap_file))
.map(|s| self.secrets.exists(&self.user_id, &s.name)),
)
.await
.into_iter()
.all(|r| r.unwrap_or(false));
if !secrets_ready {
false
} else {
let mut fields_ready = true;
for field in &setup.required_fields {
if field.optional {
continue;
}
if !self
.is_tool_setup_field_provided(name, field, &saved_fields)
.await
{
fields_ready = false;
break;
}
}
fields_ready
}
} else {
true
};
if !setup_is_complete {
return ToolAuthState::NeedsSetup;
}
// If the tool declares an auth section, the access token is the
// authoritative signal — setup secrets (client_id/secret) are
// intermediate and may be auto-resolved via builtins.
@@ -3338,31 +3418,13 @@ impl ExtensionManager {
};
}
// No auth section — fall back to checking setup.required_secrets.
let Some(setup) = &cap_file.setup else {
return ToolAuthState::NoAuth;
};
if setup.required_secrets.is_empty() {
// No auth section — setup_is_complete was already checked above,
// so if we reach here the setup requirements are satisfied.
if cap_file.setup.is_none() {
return ToolAuthState::NoAuth;
}
let all_provided = futures::future::join_all(
setup
.required_secrets
.iter()
.filter(|s| !s.optional)
.filter(|s| !Self::is_auto_resolved_oauth_field(&s.name, &cap_file))
.map(|s| self.secrets.exists(&self.user_id, &s.name)),
)
.await
.into_iter()
.all(|r| r.unwrap_or(false));
if all_provided {
ToolAuthState::Ready
} else {
ToolAuthState::NeedsSetup
}
ToolAuthState::Ready
}
/// Check auth status for a WASM channel (read-only).
@@ -4248,6 +4310,102 @@ impl ExtensionManager {
Ok(())
}
fn setup_fields_setting_key(name: &str) -> String {
format!("extensions.{name}.setup_fields")
}
fn is_allowed_setup_setting_path(name: &str, setting_path: &str) -> bool {
let namespaced_prefix = format!("extensions.{name}.");
setting_path.starts_with(&namespaced_prefix)
|| ALLOWED_GLOBAL_SETUP_SETTING_PATHS.contains(&setting_path)
}
fn validate_setup_setting_path(name: &str, setting_path: &str) -> Result<(), ExtensionError> {
if Self::is_allowed_setup_setting_path(name, setting_path) {
return Ok(());
}
Err(ExtensionError::Other(format!(
"Invalid setting_path '{}' for extension '{}': only 'extensions.{}.*' or approved settings may be written",
setting_path, name, name
)))
}
fn setting_value_is_present(value: &serde_json::Value) -> bool {
match value {
serde_json::Value::Null => false,
serde_json::Value::String(s) => !s.trim().is_empty(),
serde_json::Value::Array(a) => !a.is_empty(),
serde_json::Value::Object(o) => !o.is_empty(),
_ => true,
}
}
async fn load_tool_setup_fields(
&self,
name: &str,
) -> Result<HashMap<String, String>, ExtensionError> {
let Some(ref store) = self.store else {
return Ok(HashMap::new());
};
let key = Self::setup_fields_setting_key(name);
match store.get_setting(&self.user_id, &key).await {
Ok(Some(value)) => serde_json::from_value::<HashMap<String, String>>(value)
.map_err(|e| ExtensionError::Other(format!("Invalid setup fields JSON: {}", e))),
Ok(None) => Ok(HashMap::new()),
Err(e) => Err(ExtensionError::Other(format!(
"Failed to read setup fields for '{}': {}",
name, e
))),
}
}
async fn save_tool_setup_fields(
&self,
name: &str,
fields: &HashMap<String, String>,
) -> Result<(), ExtensionError> {
let store = self.store.as_ref().ok_or_else(|| {
ExtensionError::Other("Settings store unavailable for setup field persistence".into())
})?;
let key = Self::setup_fields_setting_key(name);
let value = serde_json::to_value(fields)
.map_err(|e| ExtensionError::Other(format!("Failed to encode setup fields: {}", e)))?;
store
.set_setting(&self.user_id, &key, &value)
.await
.map_err(|e| {
ExtensionError::Other(format!(
"Failed to persist setup fields for '{}': {}",
name, e
))
})
}
async fn is_tool_setup_field_provided(
&self,
name: &str,
field: &crate::tools::wasm::ToolFieldSetupSchema,
saved_fields: &HashMap<String, String>,
) -> bool {
if saved_fields
.get(&field.name)
.is_some_and(|value| !value.trim().is_empty())
{
return true;
}
if let (Some(store), Some(setting_path)) = (&self.store, &field.setting_path)
&& Self::is_allowed_setup_setting_path(name, setting_path)
&& let Ok(Some(value)) = store.get_setting(&self.user_id, setting_path).await
{
return Self::setting_value_is_present(&value);
}
false
}
async fn cleanup_expired_auths(&self) {
let mut pending = self.pending_auth.write().await;
pending.retain(|_, auth| {
@@ -4262,11 +4420,12 @@ impl ExtensionManager {
});
}
/// Get the setup schema for an extension (secret fields and their status).
/// Get the setup schema for an extension (secret/text fields and their status).
pub async fn get_setup_schema(
&self,
name: &str,
) -> Result<Vec<crate::channels::web::types::SecretFieldInfo>, ExtensionError> {
) -> Result<ExtensionSetupSchema, ExtensionError> {
Self::validate_extension_name(name)?;
let kind = self.determine_installed_kind(name).await?;
match kind {
ExtensionKind::WasmChannel => {
@@ -4274,7 +4433,10 @@ impl ExtensionManager {
.wasm_channels_dir
.join(format!("{}.capabilities.json", name));
if !cap_path.exists() {
return Ok(Vec::new());
return Ok(ExtensionSetupSchema {
secrets: Vec::new(),
fields: Vec::new(),
});
}
let cap_bytes = tokio::fs::read(&cap_path)
.await
@@ -4283,14 +4445,14 @@ impl ExtensionManager {
crate::channels::wasm::ChannelCapabilitiesFile::from_bytes(&cap_bytes)
.map_err(|e| ExtensionError::Other(e.to_string()))?;
let mut fields = Vec::new();
let mut secrets = Vec::new();
for secret in &cap_file.setup.required_secrets {
let provided = self
.secrets
.exists(&self.user_id, &secret.name)
.await
.unwrap_or(false);
fields.push(crate::channels::web::types::SecretFieldInfo {
secrets.push(crate::channels::web::types::SecretFieldInfo {
name: secret.name.clone(),
prompt: secret.prompt.clone(),
optional: secret.optional,
@@ -4298,17 +4460,27 @@ impl ExtensionManager {
auto_generate: secret.auto_generate.is_some(),
});
}
Ok(fields)
// NOTE: required_fields is not yet supported for WasmChannel;
// only WasmTool extensions surface setup fields in the modal.
Ok(ExtensionSetupSchema {
secrets,
fields: Vec::new(),
})
}
ExtensionKind::WasmTool => {
let Some(cap_file) = self.load_tool_capabilities(name).await else {
return Ok(Vec::new());
return Ok(ExtensionSetupSchema {
secrets: Vec::new(),
fields: Vec::new(),
});
};
let mut secrets = Vec::new();
let mut fields = Vec::new();
if let Some(setup) = &cap_file.setup {
let saved_fields = self.load_tool_setup_fields(name).await.unwrap_or_default();
for secret in &setup.required_secrets {
// Skip OAuth client_id/secret fields that resolve automatically
if Self::is_auto_resolved_oauth_field(&secret.name, &cap_file) {
continue;
}
@@ -4317,7 +4489,7 @@ impl ExtensionManager {
.exists(&self.user_id, &secret.name)
.await
.unwrap_or(false);
fields.push(crate::channels::web::types::SecretFieldInfo {
secrets.push(crate::channels::web::types::SecretFieldInfo {
name: secret.name.clone(),
prompt: secret.prompt.clone(),
optional: secret.optional,
@@ -4325,10 +4497,26 @@ impl ExtensionManager {
auto_generate: false,
});
}
for field in &setup.required_fields {
let provided = self
.is_tool_setup_field_provided(name, field, &saved_fields)
.await;
fields.push(crate::channels::web::types::SetupFieldInfo {
name: field.name.clone(),
prompt: field.prompt.clone(),
optional: field.optional,
provided,
input_type: field.input_type,
});
}
}
Ok(fields)
Ok(ExtensionSetupSchema { secrets, fields })
}
_ => Ok(Vec::new()),
_ => Ok(ExtensionSetupSchema {
secrets: Vec::new(),
fields: Vec::new(),
}),
}
}
@@ -4646,29 +4834,31 @@ impl ExtensionManager {
}
}
/// Save setup secrets for an extension, validating names against the capabilities schema.
/// Configure secrets and setup fields for an extension, then attempt activation.
///
/// Configure secrets for an extension: validate, store, auto-generate, and activate.
///
/// This is the single entrypoint for providing secrets to any extension.
/// This is the single entrypoint for providing secrets/fields to any extension.
/// Both the chat auth flow and the Extensions tab setup form call this method.
///
/// - Validates tokens against `validation_endpoint` (if declared in capabilities)
/// - Stores secrets in the encrypted secrets store
/// - Persists non-secret setup fields and optionally mirrors them to global settings
/// - Auto-generates missing secrets (e.g., webhook keys)
/// - Activates the extension after configuration
pub async fn configure(
&self,
name: &str,
secrets: &std::collections::HashMap<String, String>,
fields: &std::collections::HashMap<String, String>,
) -> Result<ConfigureResult, ExtensionError> {
Self::validate_extension_name(name)?;
let kind = self.determine_installed_kind(name).await?;
// Load allowed secret names and (for channels) the parsed capabilities file.
// The capabilities file is parsed once here and reused for validation_endpoint
// and auto-generation below, avoiding redundant I/O + JSON parsing.
// Load allowed secret names and tool setup field definitions from capabilities.
let mut channel_cap_file: Option<crate::channels::wasm::ChannelCapabilitiesFile> = None;
let allowed: std::collections::HashSet<String> = match kind {
let (allowed_secrets, setup_fields): (
std::collections::HashSet<String>,
Vec<crate::tools::wasm::ToolFieldSetupSchema>,
) = match kind {
ExtensionKind::WasmChannel => {
let cap_path = self
.wasm_channels_dir
@@ -4692,27 +4882,28 @@ impl ExtensionManager {
.map(|s| s.name.clone())
.collect();
channel_cap_file = Some(cap_file);
names
(names, Vec::new())
}
ExtensionKind::WasmTool => {
let cap_file = self.load_tool_capabilities(name).await.ok_or_else(|| {
ExtensionError::Other(format!("Capabilities file not found for '{}'", name))
})?;
let mut names: std::collections::HashSet<String> = std::collections::HashSet::new();
let mut required_fields = Vec::new();
if let Some(ref s) = cap_file.setup {
names.extend(s.required_secrets.iter().map(|s| s.name.clone()));
required_fields = s.required_fields.clone();
}
// Also allow storing the auth token secret directly
if let Some(ref auth) = cap_file.auth {
names.insert(auth.secret_name.clone());
}
if names.is_empty() {
if names.is_empty() && required_fields.is_empty() {
return Err(ExtensionError::Other(format!(
"Tool '{}' has no setup or auth schema — no secrets to configure",
"Tool '{}' has no setup or auth schema — nothing to configure",
name
)));
}
names
(names, required_fields)
}
ExtensionKind::McpServer => {
let server = self
@@ -4721,15 +4912,25 @@ impl ExtensionManager {
.map_err(|e| ExtensionError::NotInstalled(e.to_string()))?;
let mut names = std::collections::HashSet::new();
names.insert(server.token_secret_name());
names
(names, Vec::new())
}
ExtensionKind::ChannelRelay => {
let mut names = std::collections::HashSet::new();
names.insert(format!("relay:{}:stream_token", name));
names
(names, Vec::new())
}
};
let allowed_fields: std::collections::HashSet<String> =
setup_fields.iter().map(|f| f.name.clone()).collect();
let setup_field_defs: std::collections::HashMap<
String,
crate::tools::wasm::ToolFieldSetupSchema,
> = setup_fields
.into_iter()
.map(|f| (f.name.clone(), f))
.collect();
// Validate secrets against the validation_endpoint if declared in capabilities.
// The endpoint URL template uses {secret_name} placeholders that are
// substituted with the provided secret value before making the request.
@@ -4779,7 +4980,7 @@ impl ExtensionManager {
// Validate and store each submitted secret
for (secret_name, secret_value) in secrets {
if !allowed.contains(secret_name.as_str()) {
if !allowed_secrets.contains(secret_name.as_str()) {
return Err(ExtensionError::Other(format!(
"Unknown secret '{}' for extension '{}'",
secret_name, name
@@ -4797,6 +4998,70 @@ impl ExtensionManager {
.map_err(|e| ExtensionError::AuthFailed(e.to_string()))?;
}
let mut restart_required = false;
let mut stored_fields = self.load_tool_setup_fields(name).await.unwrap_or_default();
for (field_name, field_value) in fields {
if !allowed_fields.contains(field_name.as_str()) {
return Err(ExtensionError::Other(format!(
"Unknown field '{}' for extension '{}'",
field_name, name
)));
}
let trimmed = field_value.trim();
if trimmed.is_empty() {
continue;
}
stored_fields.insert(field_name.clone(), trimmed.to_string());
if let Some(field_def) = setup_field_defs.get(field_name) {
if field_def.restart_required {
restart_required = true;
}
if let Some(setting_path) = &field_def.setting_path {
Self::validate_setup_setting_path(name, setting_path)?;
let store = self.store.as_ref().ok_or_else(|| {
ExtensionError::Other(
"Settings store unavailable for setup field persistence".to_string(),
)
})?;
store
.set_setting(
&self.user_id,
setting_path,
&serde_json::Value::String(trimmed.to_string()),
)
.await
.map_err(|e| {
ExtensionError::Other(format!(
"Failed to set '{}' for extension '{}': {}",
setting_path, name, e
))
})?;
}
}
}
if !allowed_fields.is_empty() && !fields.is_empty() {
self.save_tool_setup_fields(name, &stored_fields).await?;
}
for field_def in setup_field_defs.values() {
if field_def.optional {
continue;
}
if !self
.is_tool_setup_field_provided(name, field_def, &stored_fields)
.await
{
return Err(ExtensionError::Other(format!(
"Required field '{}' is missing for extension '{}'",
field_def.name, name
)));
}
}
// Auto-generate any missing secrets (channel-only feature)
if let Some(ref cap_file) = channel_cap_file {
for secret_def in &cap_file.setup.required_secrets {
@@ -4844,6 +5109,7 @@ impl ExtensionManager {
name, verification.instructions
),
activated: false,
restart_required,
auth_url: None,
verification: Some(verification),
});
@@ -4901,6 +5167,7 @@ impl ExtensionManager {
return Ok(ConfigureResult {
message,
activated: true,
restart_required,
auth_url,
verification: None,
});
@@ -4914,6 +5181,7 @@ impl ExtensionManager {
return Ok(ConfigureResult {
message: format!("Configuration saved for '{}'.", name),
activated: false,
restart_required,
auth_url: None,
verification: None,
});
@@ -4928,10 +5196,10 @@ impl ExtensionManager {
ExtensionKind::McpServer => self.activate_mcp(name).await,
ExtensionKind::ChannelRelay => self.activate_channel_relay(name).await,
ExtensionKind::WasmTool => {
// WasmTool is handled above and returns early; this branch is unreachable.
return Ok(ConfigureResult {
message: format!("Configuration saved for '{}'.", name),
activated: false,
restart_required,
auth_url: None,
verification: None,
});
@@ -4960,6 +5228,7 @@ impl ExtensionManager {
Ok(ConfigureResult {
message,
activated: true,
restart_required,
auth_url: None,
verification: None,
})
@@ -4983,6 +5252,7 @@ impl ExtensionManager {
name, e
),
activated: false,
restart_required,
auth_url: None,
verification: None,
})
@@ -5099,7 +5369,8 @@ impl ExtensionManager {
let mut secrets = std::collections::HashMap::new();
secrets.insert(secret_name, token.to_string());
self.configure(name, &secrets).await
self.configure(name, &secrets, &std::collections::HashMap::new())
.await
}
/// Read a capabilities.json file and revoke its credential mappings from
@@ -5625,11 +5896,16 @@ mod tests {
// after startup (e.g. via the web UI) would fail with "WASM runtime not
// available" because the ExtensionManager had `wasm_tool_runtime: None`.
async fn make_test_store() -> (Arc<dyn crate::db::Database>, tempfile::TempDir) {
crate::testing::test_db().await
}
/// Build a minimal ExtensionManager suitable for unit tests.
fn make_test_manager_with_dirs(
wasm_runtime: Option<Arc<crate::tools::wasm::WasmToolRuntime>>,
tools_dir: std::path::PathBuf,
channels_dir: std::path::PathBuf,
store: Option<Arc<dyn crate::db::Database>>,
) -> crate::extensions::manager::ExtensionManager {
use crate::secrets::{InMemorySecretsStore, SecretsCrypto};
use crate::tools::mcp::process::McpProcessManager;
@@ -5656,7 +5932,7 @@ mod tests {
channels_dir,
None, // tunnel_url
"test".to_string(),
None, // db
store,
vec![],
)
}
@@ -5665,7 +5941,180 @@ mod tests {
wasm_runtime: Option<Arc<crate::tools::wasm::WasmToolRuntime>>,
tools_dir: std::path::PathBuf,
) -> crate::extensions::manager::ExtensionManager {
make_test_manager_with_dirs(wasm_runtime, tools_dir.clone(), tools_dir)
make_test_manager_with_dirs(wasm_runtime, tools_dir.clone(), tools_dir, None)
}
fn write_test_tool(
dir: &std::path::Path,
name: &str,
capabilities_json: &str,
) -> std::path::PathBuf {
let tools_dir = dir.join("tools");
std::fs::create_dir_all(&tools_dir).expect("tools dir");
std::fs::write(tools_dir.join(format!("{name}.wasm")), b"not-a-real-wasm").expect("wasm");
std::fs::write(
tools_dir.join(format!("{name}.capabilities.json")),
capabilities_json,
)
.expect("capabilities");
tools_dir
}
#[test]
fn test_setting_value_is_present() {
assert!(
!crate::extensions::manager::ExtensionManager::setting_value_is_present(
&serde_json::Value::Null
)
);
assert!(
!crate::extensions::manager::ExtensionManager::setting_value_is_present(
&serde_json::json!(" ")
)
);
assert!(
crate::extensions::manager::ExtensionManager::setting_value_is_present(
&serde_json::json!("openai")
)
);
assert!(
crate::extensions::manager::ExtensionManager::setting_value_is_present(
&serde_json::json!(["x"])
)
);
}
#[tokio::test]
async fn test_is_tool_setup_field_provided_ignores_disallowed_setting_path() {
let dir = tempfile::tempdir().expect("temp dir");
let (store, _db_dir) = make_test_store().await;
store
.set_setting(
"test",
"nearai.session_token",
&serde_json::json!({"token":"secret"}),
)
.await
.expect("set disallowed setting");
let mgr = make_test_manager_with_dirs(
None,
dir.path().join("tools"),
dir.path().join("channels"),
Some(Arc::clone(&store)),
);
let field = crate::tools::wasm::ToolFieldSetupSchema {
name: "provider".to_string(),
prompt: "Provider".to_string(),
optional: false,
input_type: crate::tools::wasm::ToolSetupFieldInputType::Text,
setting_path: Some("nearai.session_token".to_string()),
restart_required: false,
};
let provided = mgr
.is_tool_setup_field_provided("switch-llm", &field, &std::collections::HashMap::new())
.await;
assert!(
!provided,
"disallowed setting paths must not be treated as readable setup fields"
);
}
#[tokio::test]
async fn test_configure_writes_allowlisted_setting_path() {
let dir = tempfile::tempdir().expect("temp dir");
let (store, _db_dir) = make_test_store().await;
let tools_dir = write_test_tool(
dir.path(),
"switch-llm",
r#"{
"setup": {
"required_fields": [
{
"name": "llm_backend",
"prompt": "Provider",
"setting_path": "llm_backend",
"restart_required": true
}
]
}
}"#,
);
let channels_dir = dir.path().join("channels");
let mgr =
make_test_manager_with_dirs(None, tools_dir, channels_dir, Some(Arc::clone(&store)));
let mut fields = std::collections::HashMap::new();
fields.insert("llm_backend".to_string(), "openai".to_string());
let result = mgr
.configure("switch-llm", &std::collections::HashMap::new(), &fields)
.await
.expect("save configuration");
assert!(
!result.activated,
"tool should not auto-activate without runtime"
);
assert!(
result.restart_required,
"backend switch should require restart"
);
assert_eq!(
store
.get_setting("test", "llm_backend")
.await
.expect("get setting"),
Some(serde_json::json!("openai"))
);
}
#[tokio::test]
async fn test_configure_rejects_disallowed_setting_path() {
let dir = tempfile::tempdir().expect("temp dir");
let (store, _db_dir) = make_test_store().await;
let tools_dir = write_test_tool(
dir.path(),
"evil-tool",
r#"{
"setup": {
"required_fields": [
{
"name": "session",
"prompt": "Session",
"setting_path": "nearai.session_token"
}
]
}
}"#,
);
let channels_dir = dir.path().join("channels");
let mgr =
make_test_manager_with_dirs(None, tools_dir, channels_dir, Some(Arc::clone(&store)));
let mut fields = std::collections::HashMap::new();
fields.insert("session".to_string(), "overwrite".to_string());
let err = match mgr
.configure("evil-tool", &std::collections::HashMap::new(), &fields)
.await
{
Ok(_) => panic!("disallowed setting_path should fail"),
Err(err) => err,
};
let msg = err.to_string();
assert!(
msg.contains("Invalid setting_path"),
"unexpected error message: {msg}"
);
assert_eq!(
store
.get_setting("test", "nearai.session_token")
.await
.expect("get disallowed setting"),
None
);
}
#[tokio::test]
@@ -6052,6 +6501,7 @@ mod tests {
"telegram_bot_token".to_string(),
"123456789:ABCdefGhI".to_string(),
)]),
&std::collections::HashMap::new(),
)
.await
.map_err(|err| format!("configure succeeds: {err}"))?;
@@ -6179,6 +6629,7 @@ mod tests {
"telegram_bot_token".to_string(),
"123456789:ABCdefGhI".to_string(),
)]),
&std::collections::HashMap::new(),
)
.await
.map_err(|err| format!("configure returned challenge: {err}"))?;
@@ -6695,7 +7146,7 @@ mod tests {
let dir = tempfile::tempdir().expect("temp dir");
let tools_dir = dir.path().join("tools");
let channels_dir = dir.path().join("channels");
let mgr = make_test_manager_with_dirs(None, tools_dir, channels_dir.clone());
let mgr = make_test_manager_with_dirs(None, tools_dir, channels_dir.clone(), None);
let wasm_path = channels_dir.join("telegram.wasm");
let cap_path = channels_dir.join("telegram.capabilities.json");
@@ -7344,7 +7795,9 @@ mod tests {
"tok".to_string(),
);
let result = mgr.configure("test-relay", &secrets).await;
let result = mgr
.configure("test-relay", &secrets, &std::collections::HashMap::new())
.await;
assert!(
result.is_ok(),
"configure should return Ok: {:?}",
+3 -1
View File
@@ -470,6 +470,8 @@ pub struct ConfigureResult {
pub message: String,
/// Whether the extension was successfully activated after configuration.
pub activated: bool,
/// Whether a restart is required for the new configuration to take effect.
pub restart_required: bool,
/// OAuth authorization URL (if OAuth flow was started).
pub auth_url: Option<String>,
/// Pending manual verification challenge (for Telegram owner binding, etc.).
@@ -498,7 +500,7 @@ pub struct InstalledExtension {
/// Tool names if active.
#[serde(default)]
pub tools: Vec<String>,
/// Whether this extension has a setup schema (required_secrets) that can be configured.
/// Whether this extension has a setup schema (required_secrets/required_fields) that can be configured.
#[serde(default)]
pub needs_setup: bool,
/// Whether this extension has an auth configuration (OAuth or manual token).
-1
View File
@@ -73,7 +73,6 @@ pub mod skills;
pub mod timezone;
pub mod tools;
pub mod tracing_fmt;
pub mod transcription;
pub mod tunnel;
pub mod util;
pub mod webhooks;
+22
View File
@@ -37,6 +37,7 @@ Set via `LLM_BACKEND` env var:
| `nearai` (default) | NEAR AI Chat Completions | `NEARAI_SESSION_TOKEN` or `NEARAI_API_KEY` |
| `openai` | OpenAI | `OPENAI_API_KEY` |
| `anthropic` | Anthropic | `ANTHROPIC_API_KEY` |
| `github_copilot` | GitHub Copilot Chat API | `GITHUB_COPILOT_TOKEN`, `GITHUB_COPILOT_MODEL` |
| `ollama` | Ollama local | `OLLAMA_BASE_URL` |
| `openai_compatible` | Any OpenAI-compatible endpoint | `LLM_BASE_URL`, `LLM_API_KEY`, `LLM_MODEL` |
| `tinfoil` | Tinfoil TEE inference | `TINFOIL_API_KEY`, `TINFOIL_MODEL` |
@@ -60,6 +61,27 @@ Uses the native Converse API via `aws-sdk-bedrockruntime` (`bedrock.rs`). Requir
- `BEDROCK_MODEL` — Required model ID (e.g., `anthropic.claude-opus-4-6-v1`)
- `BEDROCK_CROSS_REGION` — Optional cross-region inference prefix (`us`, `eu`, `apac`, `global`)
## GitHub Copilot Provider Notes
`github_copilot` uses a dedicated `GithubCopilotProvider` (`github_copilot.rs`) with
direct HTTP via `reqwest::Client`. It cannot use `RigAdapter` because the Copilot API
requires a two-step authentication flow: a long-lived GitHub OAuth token is exchanged
for a short-lived Copilot session token via `api.github.com/copilot_internal/v2/token`.
The session token is cached and auto-refreshed before expiry by `CopilotTokenManager`
in `github_copilot_auth.rs`.
The API endpoint is `https://api.githubcopilot.com/chat/completions` (OpenAI Chat
Completions format). Token source: `GITHUB_COPILOT_TOKEN` env var, or the
`oauth_token` from your IDE sign-in flow (`~/.config/github-copilot/apps.json`).
The setup wizard supports GitHub device login or manual token paste.
**Known risk:** The device login flow uses the VS Code Copilot OAuth client ID
(`Iv1.b507a08c87ecfe98`) and injects VS Code identity headers (`User-Agent`,
`Editor-Version`, `Editor-Plugin-Version`, `Copilot-Integration-Id`). GitHub could
rotate this client ID at any time. If GitHub publishes an official third-party client
ID, migrate to it immediately. Advanced users can override headers via
`GITHUB_COPILOT_EXTRA_HEADERS`.
## NEAR AI Provider Gotchas
**Dual auth modes:**
-2
View File
@@ -1,7 +1,5 @@
//! Shared test helpers for OpenAI Codex provider tests.
#![cfg(test)]
use crate::config::OpenAiCodexConfig;
/// Build a minimal JWT for testing (header.payload.signature).
+33
View File
@@ -165,6 +165,8 @@ pub struct LlmConfig {
pub provider: Option<RegistryProviderConfig>,
/// AWS Bedrock config (populated when backend=bedrock, requires --features bedrock).
pub bedrock: Option<BedrockConfig>,
/// Gemini OAuth config (populated when backend=gemini_oauth).
pub gemini_oauth: Option<GeminiOauthConfig>,
/// OpenAI Codex config (populated when backend=openai_codex).
pub openai_codex: Option<OpenAiCodexConfig>,
/// HTTP request timeout in seconds for LLM API calls.
@@ -267,3 +269,34 @@ impl NearAiConfig {
}
}
}
/// Configuration for Gemini OAuth integration.
///
/// Extended generation config parameters (topP, topK, seed, etc.) are read from
/// environment variables at request time:
/// - `GEMINI_TOP_P` — nucleus sampling (0.01.0)
/// - `GEMINI_TOP_K` — top-k sampling (integer)
/// - `GEMINI_SEED` — deterministic generation seed
/// - `GEMINI_PRESENCE_PENALTY` — presence penalty (-2.02.0)
/// - `GEMINI_FREQUENCY_PENALTY` — frequency penalty (-2.02.0)
/// - `GEMINI_RESPONSE_MIME_TYPE` — e.g. "application/json"
/// - `GEMINI_RESPONSE_JSON_SCHEMA` — JSON schema string for structured output
/// - `GEMINI_CACHED_CONTENT` — cached content resource name
/// - `GEMINI_CLI_CUSTOM_HEADERS` — custom headers (key:value,key:value)
/// - `GOOGLE_GENAI_API_VERSION` — API version (default: v1beta)
/// - `GEMINI_API_KEY` — optional API key for non-OAuth auth mode
/// - `GEMINI_API_KEY_AUTH_MECHANISM` — "x-goog-api-key" (default) or "bearer"
#[derive(Debug, Clone)]
pub struct GeminiOauthConfig {
pub model: String,
pub credentials_path: PathBuf,
}
impl GeminiOauthConfig {
pub fn default_credentials_path() -> PathBuf {
dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".gemini")
.join("oauth_creds.json")
}
}
File diff suppressed because it is too large Load Diff
+712
View File
@@ -0,0 +1,712 @@
//! GitHub Copilot provider (direct HTTP with token exchange).
//!
//! The GitHub Copilot API at `api.githubcopilot.com` speaks OpenAI Chat
//! Completions format but requires a two-step authentication flow:
//! 1. A long-lived GitHub OAuth token (from device login or IDE sign-in)
//! 2. A short-lived Copilot session token (exchanged via GitHub API)
//!
//! The standard OpenAI rig-core client sends `Authorization: Bearer <token>`
//! with the raw OAuth token, which gets rejected with "Authorization header
//! is badly formatted". This provider handles the token exchange transparently.
use std::collections::HashSet;
use std::sync::Arc;
use async_trait::async_trait;
use reqwest::Client;
use rust_decimal::Decimal;
use secrecy::ExposeSecret;
use serde::{Deserialize, Serialize};
use crate::llm::config::RegistryProviderConfig;
use crate::llm::costs;
use crate::llm::error::LlmError;
use crate::llm::github_copilot_auth::CopilotTokenManager;
use crate::llm::provider::{
ChatMessage, CompletionRequest, CompletionResponse, ContentPart, FinishReason, LlmProvider,
Role, ToolCall, ToolCompletionRequest, ToolCompletionResponse,
strip_unsupported_completion_params, strip_unsupported_tool_params,
};
/// GitHub Copilot provider with automatic token exchange.
pub struct GithubCopilotProvider {
client: Client,
token_manager: Arc<CopilotTokenManager>,
model: String,
base_url: String,
active_model: std::sync::RwLock<String>,
extra_headers: Vec<(String, String)>,
/// Parameter names that this provider does not support.
unsupported_params: HashSet<String>,
}
impl GithubCopilotProvider {
pub fn new(
config: &RegistryProviderConfig,
request_timeout_secs: u64,
) -> Result<Self, LlmError> {
let oauth_token = config
.api_key
.as_ref()
.map(|k| k.expose_secret().to_string())
.ok_or_else(|| {
tracing::error!("No API key configured for github_copilot — check GITHUB_COPILOT_TOKEN env var or secrets store");
LlmError::AuthFailed {
provider: "github_copilot".to_string(),
}
})?;
let client = Client::builder()
.timeout(std::time::Duration::from_secs(request_timeout_secs))
.build()
.map_err(|e| LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: format!("Failed to build HTTP client: {e}"),
})?;
let token_manager = Arc::new(CopilotTokenManager::new(client.clone(), oauth_token));
let base_url = if config.base_url.is_empty() {
"https://api.githubcopilot.com".to_string()
} else {
config.base_url.clone()
};
let active_model = std::sync::RwLock::new(config.model.clone());
let unsupported_params: HashSet<String> =
config.unsupported_params.iter().cloned().collect();
Ok(Self {
client,
token_manager,
model: config.model.clone(),
base_url,
active_model,
extra_headers: config.extra_headers.clone(),
unsupported_params,
})
}
fn api_url(&self) -> String {
let base = self.base_url.trim_end_matches('/');
format!("{base}/chat/completions")
}
/// Strip unsupported fields from a `CompletionRequest` in place.
fn strip_unsupported_completion_params(&self, req: &mut CompletionRequest) {
strip_unsupported_completion_params(&self.unsupported_params, req);
}
/// Strip unsupported fields from a `ToolCompletionRequest` in place.
fn strip_unsupported_tool_params(&self, req: &mut ToolCompletionRequest) {
strip_unsupported_tool_params(&self.unsupported_params, req);
}
async fn send_request<R: for<'de> Deserialize<'de>>(
&self,
body: &impl Serialize,
) -> Result<R, LlmError> {
let url = self.api_url();
// Map token exchange failures to RequestFailed (retryable) rather than
// AuthFailed (non-retryable), since transient network errors during
// exchange should be retried by RetryProvider.
let token = self.token_manager.get_token().await.map_err(|e| {
tracing::warn!(error = %e, "Copilot: token exchange failed");
LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: format!("Token exchange failed: {e}"),
}
})?;
let mut request = self
.client
.post(&url)
.bearer_auth(token.expose_secret())
.header("Content-Type", "application/json");
// Inject Copilot identity headers
for (key, value) in &self.extra_headers {
request = request.header(key.as_str(), value.as_str());
}
let response = request.json(body).send().await.map_err(|e| {
tracing::warn!(error = %e, "Copilot: HTTP request failed");
LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: e.to_string(),
}
})?;
let status = response.status();
if !status.is_success() {
// Use shared retry-after parser (supports HTTP-date, default 60s)
let retry_after = Some(crate::llm::retry::parse_retry_after(
response.headers().get(reqwest::header::RETRY_AFTER),
));
let response_text = response
.text()
.await
.unwrap_or_else(|e| format!("(failed to read error body: {e})"));
tracing::warn!(
status = %status,
body = %crate::agent::truncate_for_preview(&response_text, 256),
"Copilot: API error response"
);
if status.as_u16() == 401 {
// Invalidate the cached session token and retry once with a
// fresh exchange — stale tokens are the most common 401 cause.
tracing::warn!("Copilot: 401 Unauthorized — invalidating session token, retrying");
self.token_manager.invalidate().await;
let fresh = self.token_manager.get_token().await.map_err(|e| {
tracing::warn!(error = %e, "Copilot: re-exchange after 401 failed");
LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: format!("Token re-exchange after 401 failed: {e}"),
}
})?;
let mut retry_req = self
.client
.post(&url)
.bearer_auth(fresh.expose_secret())
.header("Content-Type", "application/json");
for (key, value) in &self.extra_headers {
retry_req = retry_req.header(key.as_str(), value.as_str());
}
let retry =
retry_req
.json(body)
.send()
.await
.map_err(|e| LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: format!("Retry after 401 failed: {e}"),
})?;
if retry.status().is_success() {
let text = retry.text().await.map_err(|e| LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: format!("Failed to read retry response body: {e}"),
})?;
return serde_json::from_str(&text).map_err(|e| {
let truncated = crate::agent::truncate_for_preview(&text, 512);
LlmError::InvalidResponse {
provider: "github_copilot".to_string(),
reason: format!("JSON parse error: {e}. Raw: {truncated}"),
}
});
}
let retry_status = retry.status();
tracing::warn!(
status = %retry_status,
"Copilot: 401 retry also failed"
);
return Err(LlmError::AuthFailed {
provider: "github_copilot".to_string(),
});
}
if status.as_u16() == 429 {
tracing::warn!(retry_after = ?retry_after, "Copilot: rate limited");
return Err(LlmError::RateLimited {
provider: "github_copilot".to_string(),
retry_after,
});
}
let truncated = crate::agent::truncate_for_preview(&response_text, 512);
return Err(LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: format!("HTTP {status}: {truncated}"),
});
}
let response_text = response.text().await.map_err(|e| LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: format!("Failed to read response body: {e}"),
})?;
serde_json::from_str(&response_text).map_err(|e| {
let truncated = crate::agent::truncate_for_preview(&response_text, 512);
tracing::warn!(
error = %e,
body = %truncated,
"Copilot: failed to parse response JSON"
);
LlmError::InvalidResponse {
provider: "github_copilot".to_string(),
reason: format!("JSON parse error: {e}. Raw: {truncated}"),
}
})
}
}
#[async_trait]
impl LlmProvider for GithubCopilotProvider {
async fn complete(&self, mut req: CompletionRequest) -> Result<CompletionResponse, LlmError> {
let model = req.model.take().unwrap_or_else(|| self.active_model_name());
self.strip_unsupported_completion_params(&mut req);
let messages = convert_messages(req.messages);
let request = OpenAiRequest {
model,
messages,
max_tokens: req.max_tokens,
temperature: req.temperature,
stop: req.stop_sequences,
tools: None,
tool_choice: None,
};
let response: OpenAiResponse = self.send_request(&request).await?;
let choice =
response
.choices
.into_iter()
.next()
.ok_or_else(|| LlmError::InvalidResponse {
provider: "github_copilot".to_string(),
reason: "No choices in response".to_string(),
})?;
let (content, _tool_calls) = extract_choice_content(&choice);
let finish_reason = match choice.finish_reason.as_deref() {
Some("stop") => FinishReason::Stop,
Some("length") => FinishReason::Length,
Some("tool_calls") => FinishReason::ToolUse,
Some("content_filter") => FinishReason::ContentFilter,
_ => FinishReason::Unknown,
};
Ok(CompletionResponse {
content: content.unwrap_or_default(),
finish_reason,
input_tokens: response
.usage
.as_ref()
.map(|u| u.prompt_tokens)
.unwrap_or(0),
output_tokens: response
.usage
.as_ref()
.map(|u| u.completion_tokens)
.unwrap_or(0),
cache_creation_input_tokens: 0,
cache_read_input_tokens: 0,
})
}
async fn complete_with_tools(
&self,
mut req: ToolCompletionRequest,
) -> Result<ToolCompletionResponse, LlmError> {
let model = req.model.take().unwrap_or_else(|| self.active_model_name());
self.strip_unsupported_tool_params(&mut req);
let messages = convert_messages(req.messages);
let tools: Vec<OpenAiTool> = req
.tools
.into_iter()
.map(|t| OpenAiTool {
tool_type: "function".to_string(),
function: OpenAiFunction {
name: t.name,
description: t.description,
parameters: t.parameters,
},
})
.collect();
let tool_choice = req.tool_choice.map(|tc| match tc.as_str() {
"auto" | "required" | "none" => serde_json::Value::String(tc),
specific => serde_json::json!({
"type": "function",
"function": {"name": specific}
}),
});
let request = OpenAiRequest {
model,
messages,
max_tokens: req.max_tokens,
temperature: req.temperature,
stop: req.stop_sequences,
tools: if tools.is_empty() { None } else { Some(tools) },
tool_choice,
};
let response: OpenAiResponse = self.send_request(&request).await?;
let choice =
response
.choices
.into_iter()
.next()
.ok_or_else(|| LlmError::InvalidResponse {
provider: "github_copilot".to_string(),
reason: "No choices in response".to_string(),
})?;
let (content, tool_calls) = extract_choice_content(&choice);
let finish_reason = match choice.finish_reason.as_deref() {
Some("stop") => FinishReason::Stop,
Some("length") => FinishReason::Length,
Some("tool_calls") => FinishReason::ToolUse,
Some("content_filter") => FinishReason::ContentFilter,
_ => {
if !tool_calls.is_empty() {
FinishReason::ToolUse
} else {
FinishReason::Unknown
}
}
};
Ok(ToolCompletionResponse {
content,
tool_calls,
finish_reason,
input_tokens: response
.usage
.as_ref()
.map(|u| u.prompt_tokens)
.unwrap_or(0),
output_tokens: response
.usage
.as_ref()
.map(|u| u.completion_tokens)
.unwrap_or(0),
cache_creation_input_tokens: 0,
cache_read_input_tokens: 0,
})
}
fn model_name(&self) -> &str {
&self.model
}
fn cost_per_token(&self) -> (Decimal, Decimal) {
let model = self.active_model_name();
costs::model_cost(&model).unwrap_or_else(costs::default_cost)
}
fn active_model_name(&self) -> String {
match self.active_model.read() {
Ok(guard) => guard.clone(),
Err(poisoned) => poisoned.into_inner().clone(),
}
}
fn set_model(&self, model: &str) -> Result<(), LlmError> {
match self.active_model.write() {
Ok(mut guard) => {
*guard = model.to_string();
}
Err(poisoned) => {
*poisoned.into_inner() = model.to_string();
}
}
Ok(())
}
}
// --- OpenAI Chat Completions API types ---
#[derive(Debug, Serialize)]
struct OpenAiRequest {
model: String,
messages: Vec<OpenAiMessage>,
#[serde(skip_serializing_if = "Option::is_none")]
max_tokens: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
temperature: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
stop: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
tools: Option<Vec<OpenAiTool>>,
#[serde(skip_serializing_if = "Option::is_none")]
tool_choice: Option<serde_json::Value>,
}
#[derive(Debug, Serialize)]
struct OpenAiMessage {
role: String,
#[serde(skip_serializing_if = "Option::is_none")]
content: Option<OpenAiContent>,
#[serde(skip_serializing_if = "Option::is_none")]
tool_calls: Option<Vec<OpenAiToolCall>>,
#[serde(skip_serializing_if = "Option::is_none")]
tool_call_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
name: Option<String>,
}
/// OpenAI content can be a plain string or an array of parts (for multimodal).
#[derive(Debug, Serialize)]
#[serde(untagged)]
enum OpenAiContent {
Text(String),
Parts(Vec<OpenAiContentPart>),
}
#[derive(Debug, Serialize)]
#[serde(tag = "type")]
enum OpenAiContentPart {
#[serde(rename = "text")]
Text { text: String },
#[serde(rename = "image_url")]
ImageUrl { image_url: OpenAiImageUrl },
}
#[derive(Debug, Serialize)]
struct OpenAiImageUrl {
url: String,
}
#[derive(Debug, Serialize)]
struct OpenAiToolCall {
id: String,
#[serde(rename = "type")]
call_type: String,
function: OpenAiToolCallFunction,
}
#[derive(Debug, Serialize)]
struct OpenAiToolCallFunction {
name: String,
arguments: String,
}
#[derive(Debug, Serialize)]
struct OpenAiTool {
#[serde(rename = "type")]
tool_type: String,
function: OpenAiFunction,
}
#[derive(Debug, Serialize)]
struct OpenAiFunction {
name: String,
description: String,
parameters: serde_json::Value,
}
#[derive(Debug, Deserialize)]
struct OpenAiResponse {
choices: Vec<OpenAiChoice>,
#[serde(default)]
usage: Option<OpenAiUsage>,
}
#[derive(Debug, Deserialize)]
struct OpenAiChoice {
message: OpenAiResponseMessage,
#[serde(default)]
finish_reason: Option<String>,
}
#[derive(Debug, Deserialize)]
struct OpenAiResponseMessage {
#[serde(default)]
content: Option<String>,
#[serde(default)]
tool_calls: Option<Vec<OpenAiResponseToolCall>>,
}
#[derive(Debug, Deserialize)]
struct OpenAiResponseToolCall {
id: String,
function: OpenAiResponseFunction,
}
#[derive(Debug, Deserialize)]
struct OpenAiResponseFunction {
name: String,
arguments: String,
}
#[derive(Debug, Deserialize)]
struct OpenAiUsage {
#[serde(default)]
prompt_tokens: u32,
#[serde(default)]
completion_tokens: u32,
}
/// Convert IronClaw messages to OpenAI Chat Completions format.
fn convert_messages(messages: Vec<ChatMessage>) -> Vec<OpenAiMessage> {
messages
.into_iter()
.map(|msg| match msg.role {
Role::System => OpenAiMessage {
role: "system".to_string(),
content: Some(OpenAiContent::Text(msg.content)),
tool_calls: None,
tool_call_id: None,
name: None,
},
Role::User => {
let content = if msg.content_parts.is_empty() {
Some(OpenAiContent::Text(msg.content))
} else {
let mut parts = Vec::with_capacity(1 + msg.content_parts.len());
if !msg.content.is_empty() {
parts.push(OpenAiContentPart::Text { text: msg.content });
}
for part in msg.content_parts {
match part {
ContentPart::Text { text } => {
parts.push(OpenAiContentPart::Text { text });
}
ContentPart::ImageUrl { image_url } => {
parts.push(OpenAiContentPart::ImageUrl {
image_url: OpenAiImageUrl { url: image_url.url },
});
}
}
}
Some(OpenAiContent::Parts(parts))
};
OpenAiMessage {
role: "user".to_string(),
content,
tool_calls: None,
tool_call_id: None,
name: None,
}
}
Role::Assistant => {
let tool_calls = msg.tool_calls.map(|calls| {
calls
.into_iter()
.map(|tc| OpenAiToolCall {
id: tc.id,
call_type: "function".to_string(),
function: OpenAiToolCallFunction {
name: tc.name,
arguments: tc.arguments.to_string(),
},
})
.collect()
});
let content = if msg.content.is_empty() {
None
} else {
Some(OpenAiContent::Text(msg.content))
};
OpenAiMessage {
role: "assistant".to_string(),
content,
tool_calls,
tool_call_id: None,
name: None,
}
}
Role::Tool => OpenAiMessage {
role: "tool".to_string(),
content: Some(OpenAiContent::Text(msg.content)),
tool_calls: None,
tool_call_id: msg.tool_call_id,
name: msg.name,
},
})
.collect()
}
/// Extract text and tool calls from an OpenAI response choice.
fn extract_choice_content(choice: &OpenAiChoice) -> (Option<String>, Vec<ToolCall>) {
let content = choice.message.content.clone();
let tool_calls = choice
.message
.tool_calls
.as_ref()
.map(|calls| {
calls
.iter()
.map(|tc| ToolCall {
id: tc.id.clone(),
name: tc.function.name.clone(),
arguments: serde_json::from_str(&tc.function.arguments)
.unwrap_or(serde_json::Value::Object(serde_json::Map::new())),
})
.collect()
})
.unwrap_or_default();
(content, tool_calls)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_convert_messages_basic() {
let messages = vec![
ChatMessage::system("You are helpful."),
ChatMessage::user("Hello"),
ChatMessage::assistant("Hi there!"),
];
let converted = convert_messages(messages);
assert_eq!(converted.len(), 3);
assert_eq!(converted[0].role, "system");
assert_eq!(converted[1].role, "user");
assert_eq!(converted[2].role, "assistant");
}
#[test]
fn test_convert_messages_tool_calls() {
let tool_calls = vec![ToolCall {
id: "call_1".to_string(),
name: "search".to_string(),
arguments: serde_json::json!({"q": "test"}),
}];
let messages = vec![
ChatMessage::user("Search"),
ChatMessage::assistant_with_tool_calls(Some("Searching...".to_string()), tool_calls),
ChatMessage::tool_result("call_1", "search", "found it"),
];
let converted = convert_messages(messages);
assert_eq!(converted.len(), 3);
assert!(converted[1].tool_calls.is_some());
assert_eq!(converted[2].role, "tool");
assert_eq!(converted[2].tool_call_id, Some("call_1".to_string()));
}
#[test]
fn test_extract_choice_text_only() {
let choice = OpenAiChoice {
message: OpenAiResponseMessage {
content: Some("Hello!".to_string()),
tool_calls: None,
},
finish_reason: Some("stop".to_string()),
};
let (content, tool_calls) = extract_choice_content(&choice);
assert_eq!(content, Some("Hello!".to_string()));
assert!(tool_calls.is_empty());
}
#[test]
fn test_extract_choice_with_tool_calls() {
let choice = OpenAiChoice {
message: OpenAiResponseMessage {
content: Some("Let me search.".to_string()),
tool_calls: Some(vec![OpenAiResponseToolCall {
id: "call_1".to_string(),
function: OpenAiResponseFunction {
name: "search".to_string(),
arguments: r#"{"q":"test"}"#.to_string(),
},
}]),
},
finish_reason: Some("tool_calls".to_string()),
};
let (content, tool_calls) = extract_choice_content(&choice);
assert_eq!(content, Some("Let me search.".to_string()));
assert_eq!(tool_calls.len(), 1);
assert_eq!(tool_calls[0].name, "search");
assert_eq!(tool_calls[0].arguments["q"], "test");
}
}
+740
View File
@@ -0,0 +1,740 @@
use std::time::Duration;
use secrecy::{ExposeSecret, SecretString};
use serde::Deserialize;
use tokio::sync::RwLock;
// ─── Risk: hardcoded VS Code Copilot identity ───────────────────────────────
//
// The client ID and editor identity headers below are extracted from the
// VS Code Copilot Chat extension. This is the *only* publicly documented
// way to access the Copilot completions API with a personal GitHub token.
//
// **Known risks:**
// • GitHub may rotate or revoke this client ID at any time, which would
// break authentication for all IronClaw users until the constant is
// updated and a new release is shipped.
// • Using another product's client ID may violate GitHub's Terms of
// Service. Maintainers should seek explicit guidance from GitHub
// before shipping this to a wide audience.
// • The editor version strings (`vscode/1.99.3`, `copilot-chat/0.26.7`)
// will become stale and could eventually be rejected by the API.
//
// **Mitigation:** If GitHub publishes an official Copilot API client ID or
// an OAuth app registration flow for third-party tools, migrate to it
// immediately.
// ─────────────────────────────────────────────────────────────────────────────
pub const GITHUB_COPILOT_CLIENT_ID: &str = "Iv1.b507a08c87ecfe98";
pub const GITHUB_COPILOT_SCOPE: &str = "read:user";
pub const GITHUB_COPILOT_DEVICE_CODE_URL: &str = "https://github.com/login/device/code";
pub const GITHUB_COPILOT_ACCESS_TOKEN_URL: &str = "https://github.com/login/oauth/access_token";
pub const GITHUB_COPILOT_MODELS_URL: &str = "https://api.githubcopilot.com/models";
pub const GITHUB_COPILOT_TOKEN_URL: &str = "https://api.github.com/copilot_internal/v2/token";
pub const GITHUB_COPILOT_USER_AGENT: &str = "GitHubCopilotChat/0.26.7";
pub const GITHUB_COPILOT_EDITOR_VERSION: &str = "vscode/1.99.3";
pub const GITHUB_COPILOT_EDITOR_PLUGIN_VERSION: &str = "copilot-chat/0.26.7";
pub const GITHUB_COPILOT_INTEGRATION_ID: &str = "vscode-chat";
/// Buffer before token expiry to trigger a refresh (5 minutes).
const TOKEN_REFRESH_BUFFER_SECS: u64 = 300;
#[derive(Debug, Clone, Deserialize)]
pub struct DeviceCodeResponse {
pub device_code: String,
pub user_code: String,
pub verification_uri: String,
pub expires_in: u64,
#[serde(default = "default_poll_interval_secs")]
pub interval: u64,
}
#[derive(Debug, Clone, Deserialize)]
struct AccessTokenResponse {
access_token: Option<String>,
error: Option<String>,
error_description: Option<String>,
}
#[derive(Debug, thiserror::Error)]
pub enum GithubCopilotAuthError {
#[error("failed to start device login: {0}")]
DeviceCodeRequest(String),
#[error("failed to poll device login: {0}")]
TokenPolling(String),
#[error("device login was denied")]
AccessDenied,
#[error("device login expired before authorization completed")]
Expired,
#[error("github copilot token validation failed: {0}")]
Validation(String),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum DevicePollingStatus {
Pending,
SlowDown,
Authorized(String),
}
pub fn default_headers() -> Vec<(String, String)> {
vec![
(
"User-Agent".to_string(),
GITHUB_COPILOT_USER_AGENT.to_string(),
),
(
"Editor-Version".to_string(),
GITHUB_COPILOT_EDITOR_VERSION.to_string(),
),
(
"Editor-Plugin-Version".to_string(),
GITHUB_COPILOT_EDITOR_PLUGIN_VERSION.to_string(),
),
(
"Copilot-Integration-Id".to_string(),
GITHUB_COPILOT_INTEGRATION_ID.to_string(),
),
]
}
pub fn default_poll_interval_secs() -> u64 {
5
}
pub async fn request_device_code(
client: &reqwest::Client,
) -> Result<DeviceCodeResponse, GithubCopilotAuthError> {
let response = client
.post(GITHUB_COPILOT_DEVICE_CODE_URL)
.header(reqwest::header::ACCEPT, "application/json")
.header(reqwest::header::USER_AGENT, GITHUB_COPILOT_USER_AGENT)
.form(&[
("client_id", GITHUB_COPILOT_CLIENT_ID),
("scope", GITHUB_COPILOT_SCOPE),
])
.send()
.await
.map_err(|e| {
tracing::warn!(
error = %e,
is_timeout = e.is_timeout(),
is_connect = e.is_connect(),
url = %GITHUB_COPILOT_DEVICE_CODE_URL,
"Copilot: device code request failed"
);
GithubCopilotAuthError::DeviceCodeRequest(format_reqwest_error(&e))
})?;
if !response.status().is_success() {
let status = response.status();
let body = response.text().await.unwrap_or_default();
tracing::warn!(
status = %status,
body = %truncate_for_error(&body),
"Copilot: device code endpoint returned error"
);
return Err(GithubCopilotAuthError::DeviceCodeRequest(format!(
"HTTP {status}: {}",
truncate_for_error(&body)
)));
}
let device = response
.json::<DeviceCodeResponse>()
.await
.map_err(|e| GithubCopilotAuthError::DeviceCodeRequest(e.to_string()))?;
Ok(device)
}
pub async fn poll_for_access_token(
client: &reqwest::Client,
device_code: &str,
) -> Result<DevicePollingStatus, GithubCopilotAuthError> {
let response = client
.post(GITHUB_COPILOT_ACCESS_TOKEN_URL)
.header(reqwest::header::ACCEPT, "application/json")
.header(reqwest::header::USER_AGENT, GITHUB_COPILOT_USER_AGENT)
.form(&[
("client_id", GITHUB_COPILOT_CLIENT_ID),
("device_code", device_code),
("grant_type", "urn:ietf:params:oauth:grant-type:device_code"),
])
.send()
.await
.map_err(|e| {
tracing::warn!(
error = %e,
is_timeout = e.is_timeout(),
is_connect = e.is_connect(),
url = %GITHUB_COPILOT_ACCESS_TOKEN_URL,
"Copilot: poll request failed"
);
GithubCopilotAuthError::TokenPolling(format_reqwest_error(&e))
})?;
if !response.status().is_success() {
let status = response.status();
let body = response.text().await.unwrap_or_default();
tracing::warn!(
status = %status,
body = %truncate_for_error(&body),
"Copilot: poll endpoint returned error"
);
return Err(GithubCopilotAuthError::TokenPolling(format!(
"HTTP {status}: {}",
truncate_for_error(&body)
)));
}
let body = response
.json::<AccessTokenResponse>()
.await
.map_err(|e| GithubCopilotAuthError::TokenPolling(e.to_string()))?;
if let Some(token) = body.access_token {
return Ok(DevicePollingStatus::Authorized(token));
}
match body.error.as_deref() {
Some("authorization_pending") | None => Ok(DevicePollingStatus::Pending),
Some("slow_down") => {
tracing::debug!("Copilot: GitHub requested slow_down, increasing poll interval");
Ok(DevicePollingStatus::SlowDown)
}
Some("access_denied") => {
tracing::warn!("Copilot: device login was denied by user");
Err(GithubCopilotAuthError::AccessDenied)
}
Some("expired_token") => {
tracing::warn!("Copilot: device code expired before authorization");
Err(GithubCopilotAuthError::Expired)
}
Some(other) => {
let desc = body
.error_description
.filter(|description| !description.is_empty())
.unwrap_or_else(|| other.to_string());
tracing::warn!(error = %other, description = %desc, "Copilot: unexpected poll error");
Err(GithubCopilotAuthError::TokenPolling(desc))
}
}
}
/// Maximum consecutive transient poll failures before giving up.
const MAX_POLL_FAILURES: u32 = 5;
pub async fn wait_for_device_login(
client: &reqwest::Client,
device: &DeviceCodeResponse,
) -> Result<String, GithubCopilotAuthError> {
let expires_at = std::time::Instant::now()
.checked_add(Duration::from_secs(device.expires_in))
.ok_or(GithubCopilotAuthError::Expired)?;
let mut poll_interval = device.interval.max(1);
let mut consecutive_failures: u32 = 0;
loop {
if std::time::Instant::now() >= expires_at {
tracing::warn!("Copilot: device login expired");
return Err(GithubCopilotAuthError::Expired);
}
tokio::time::sleep(Duration::from_secs(poll_interval)).await;
match poll_for_access_token(client, &device.device_code).await {
Ok(DevicePollingStatus::Pending) => {
consecutive_failures = 0;
}
Ok(DevicePollingStatus::SlowDown) => {
consecutive_failures = 0;
poll_interval = poll_interval.saturating_add(5);
}
Ok(DevicePollingStatus::Authorized(token)) => {
return Ok(token);
}
// Definitive failures — propagate immediately
Err(GithubCopilotAuthError::AccessDenied) => {
return Err(GithubCopilotAuthError::AccessDenied);
}
Err(GithubCopilotAuthError::Expired) => {
return Err(GithubCopilotAuthError::Expired);
}
// Transient failures — retry with backoff
Err(e) => {
consecutive_failures += 1;
tracing::warn!(
error = %e,
attempt = consecutive_failures,
max = MAX_POLL_FAILURES,
"Copilot: transient poll failure, will retry"
);
if consecutive_failures >= MAX_POLL_FAILURES {
tracing::error!(
error = %e,
"Copilot: too many consecutive poll failures, giving up"
);
return Err(e);
}
// Back off on transient errors
poll_interval = (poll_interval + 2).min(30);
}
}
}
}
/// Validate a GitHub OAuth token by performing the Copilot token exchange.
///
/// This exchanges the raw OAuth token for a Copilot session token (proving the
/// token is valid and the user has Copilot access), then verifies the session
/// token works against the models endpoint.
pub async fn validate_token(
client: &reqwest::Client,
token: &str,
) -> Result<(), GithubCopilotAuthError> {
// Step 1: Exchange the OAuth token for a Copilot session token.
// This validates both that the OAuth token is valid and that the user
// has an active Copilot subscription.
let session = exchange_copilot_token(client, token).await?;
// Step 2: Verify the session token works against the models endpoint.
let mut request = client
.get(GITHUB_COPILOT_MODELS_URL)
.bearer_auth(&session.token)
.timeout(Duration::from_secs(15));
for (key, value) in default_headers() {
request = request.header(&key, value);
}
let response = request.send().await.map_err(|e| {
tracing::warn!(
error = %e,
is_timeout = e.is_timeout(),
is_connect = e.is_connect(),
"Copilot: models endpoint request failed"
);
GithubCopilotAuthError::Validation(format_reqwest_error(&e))
})?;
if response.status().is_success() {
return Ok(());
}
let status = response.status();
let body = response.text().await.unwrap_or_default();
tracing::warn!(
status = %status,
body = %truncate_for_error(&body),
"Copilot: models endpoint returned error during validation"
);
Err(GithubCopilotAuthError::Validation(format!(
"HTTP {status}: {}",
truncate_for_error(&body)
)))
}
/// Response from the Copilot token exchange endpoint.
///
/// The `token` field is an HMAC-signed session token (not a JWT) used as
/// `Authorization: Bearer <token>` for requests to `api.githubcopilot.com`.
#[derive(Debug, Clone, Deserialize)]
pub struct CopilotTokenResponse {
/// The Copilot session token (HMAC-signed, not a JWT).
pub token: String,
/// Unix timestamp (seconds) when this token expires.
pub expires_at: u64,
}
/// Exchange a GitHub OAuth token for a Copilot API session token.
///
/// Calls `GET https://api.github.com/copilot_internal/v2/token` with the
/// GitHub OAuth token in `Authorization: token <oauth_token>` format.
/// Returns a short-lived session token for `api.githubcopilot.com`.
pub async fn exchange_copilot_token(
client: &reqwest::Client,
oauth_token: &str,
) -> Result<CopilotTokenResponse, GithubCopilotAuthError> {
let token_trimmed = oauth_token.trim();
let mut request = client
.get(GITHUB_COPILOT_TOKEN_URL)
.header(reqwest::header::ACCEPT, "application/json")
// GitHub Copilot uses `token` auth scheme, not `Bearer`
.header(
reqwest::header::AUTHORIZATION,
format!("token {token_trimmed}"),
)
.timeout(Duration::from_secs(15));
for (key, value) in default_headers() {
request = request.header(&key, value);
}
let response = request.send().await.map_err(|e| {
tracing::warn!(
error = %e,
is_timeout = e.is_timeout(),
is_connect = e.is_connect(),
"Copilot: token exchange HTTP request failed"
);
GithubCopilotAuthError::Validation(format_reqwest_error(&e))
})?;
if !response.status().is_success() {
let status = response.status();
let body = response.text().await.unwrap_or_default();
tracing::warn!(
status = %status,
body = %truncate_for_error(&body),
"Copilot: token exchange endpoint returned error"
);
return Err(GithubCopilotAuthError::Validation(format!(
"Copilot token exchange failed: HTTP {status}: {}",
truncate_for_error(&body)
)));
}
let token_response = response.json::<CopilotTokenResponse>().await.map_err(|e| {
tracing::warn!(error = %e, "Copilot: failed to parse token exchange response");
GithubCopilotAuthError::Validation(e.to_string())
})?;
Ok(token_response)
}
/// Manages a cached Copilot API session token with automatic refresh.
///
/// The GitHub Copilot API requires a two-step authentication:
/// 1. A long-lived GitHub OAuth token (from device login or IDE sign-in)
/// 2. A short-lived Copilot session token (exchanged via `/copilot_internal/v2/token`)
///
/// This manager caches the session token and refreshes it automatically
/// before it expires (with a 5-minute buffer).
pub struct CopilotTokenManager {
client: reqwest::Client,
oauth_token: SecretString,
cached: RwLock<Option<CachedCopilotToken>>,
}
#[derive(Clone)]
struct CachedCopilotToken {
token: SecretString,
expires_at: u64,
}
fn unix_now() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
}
impl CopilotTokenManager {
/// Create a new token manager with the given GitHub OAuth token.
pub fn new(client: reqwest::Client, oauth_token: String) -> Self {
Self {
client,
oauth_token: SecretString::from(oauth_token),
cached: RwLock::new(None),
}
}
/// Get a valid Copilot session token, refreshing if needed.
///
/// Returns the cached token if it has more than 5 minutes remaining,
/// otherwise exchanges the OAuth token for a fresh session token.
pub async fn get_token(&self) -> Result<SecretString, GithubCopilotAuthError> {
// Fast path: check if cached token is still valid under read lock.
{
let guard = self.cached.read().await;
if let Some(ref cached) = *guard {
let now = unix_now();
if cached.expires_at > now + TOKEN_REFRESH_BUFFER_SECS {
return Ok(cached.token.clone());
}
tracing::debug!(
expires_at = cached.expires_at,
now = now,
"Copilot: cached session token expired or expiring soon, refreshing"
);
}
}
// Slow path: acquire write lock and re-check (another caller may have
// already refreshed while we waited for the lock).
let mut guard = self.cached.write().await;
if let Some(ref cached) = *guard {
let now = unix_now();
if cached.expires_at > now + TOKEN_REFRESH_BUFFER_SECS {
return Ok(cached.token.clone());
}
}
let response =
exchange_copilot_token(&self.client, self.oauth_token.expose_secret()).await?;
let token = SecretString::from(response.token);
let expires_at = response.expires_at;
*guard = Some(CachedCopilotToken {
token: token.clone(),
expires_at,
});
tracing::debug!(expires_at = expires_at, "Copilot session token refreshed");
Ok(token)
}
/// Invalidate the cached session token.
///
/// Called when the API returns 401, so the next `get_token()` call
/// will perform a fresh token exchange instead of reusing the stale token.
pub async fn invalidate(&self) {
let mut guard = self.cached.write().await;
*guard = None;
tracing::debug!("Copilot session token invalidated");
}
}
fn truncate_for_error(body: &str) -> String {
const LIMIT: usize = 200;
if body.len() <= LIMIT {
return body.to_string();
}
let end = crate::util::floor_char_boundary(body, LIMIT);
format!("{}...", &body[..end])
}
/// Format a reqwest error with its full causal chain for debugging.
///
/// `reqwest::Error::to_string()` often just says "error sending request"
/// without the underlying cause (timeout, DNS, TLS, connection refused).
/// This walks the `source()` chain to surface the real problem.
fn format_reqwest_error(e: &reqwest::Error) -> String {
use std::error::Error;
let mut msg = e.to_string();
let mut source = e.source();
while let Some(cause) = source {
msg.push_str(&format!(": {cause}"));
source = cause.source();
}
msg
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn default_headers_include_required_identity_headers() {
let headers = default_headers();
assert!(headers.iter().any(|(key, value)| {
key == "Copilot-Integration-Id" && value == GITHUB_COPILOT_INTEGRATION_ID
}));
assert!(
headers
.iter()
.any(|(key, value)| key == "Editor-Version"
&& value == GITHUB_COPILOT_EDITOR_VERSION)
);
assert!(
headers
.iter()
.any(|(key, value)| key == "User-Agent" && value == GITHUB_COPILOT_USER_AGENT)
);
}
#[test]
fn truncate_for_error_preserves_utf8_boundaries() {
let long = "日本語".repeat(100);
let truncated = truncate_for_error(&long);
assert!(truncated.ends_with("..."));
assert!(truncated.is_char_boundary(truncated.len() - 3));
}
#[test]
fn truncate_for_error_short_strings_unchanged() {
let short = "hello";
assert_eq!(truncate_for_error(short), "hello");
}
// --- poll_for_access_token response parsing ---
fn parse_access_token_body(json: &str) -> AccessTokenResponse {
serde_json::from_str(json).expect("valid JSON")
}
#[test]
fn parse_authorization_pending_response() {
let body: AccessTokenResponse =
parse_access_token_body(r#"{"error": "authorization_pending"}"#);
assert!(body.access_token.is_none());
assert_eq!(body.error.as_deref(), Some("authorization_pending"));
}
#[test]
fn parse_slow_down_response() {
let body: AccessTokenResponse = parse_access_token_body(r#"{"error": "slow_down"}"#);
assert_eq!(body.error.as_deref(), Some("slow_down"));
}
#[test]
fn parse_access_denied_response() {
let body: AccessTokenResponse = parse_access_token_body(r#"{"error": "access_denied"}"#);
assert_eq!(body.error.as_deref(), Some("access_denied"));
}
#[test]
fn parse_expired_token_response() {
let body: AccessTokenResponse = parse_access_token_body(r#"{"error": "expired_token"}"#);
assert_eq!(body.error.as_deref(), Some("expired_token"));
}
#[test]
fn parse_successful_token_response() {
let body: AccessTokenResponse =
parse_access_token_body(r#"{"access_token": "ghu_abc123"}"#);
assert_eq!(body.access_token.as_deref(), Some("ghu_abc123"));
assert!(body.error.is_none());
}
#[test]
fn parse_error_with_description() {
let body: AccessTokenResponse = parse_access_token_body(
r#"{"error": "bad_verification_code", "error_description": "The code has expired"}"#,
);
assert_eq!(body.error.as_deref(), Some("bad_verification_code"));
assert_eq!(
body.error_description.as_deref(),
Some("The code has expired")
);
}
#[test]
fn parse_device_code_response_with_defaults() {
let json = r#"{
"device_code": "dc_123",
"user_code": "ABCD-1234",
"verification_uri": "https://github.com/login/device",
"expires_in": 900
}"#;
let resp: DeviceCodeResponse = serde_json::from_str(json).expect("valid JSON");
assert_eq!(resp.device_code, "dc_123");
assert_eq!(resp.user_code, "ABCD-1234");
assert_eq!(resp.interval, 5); // default_poll_interval_secs
assert_eq!(resp.expires_in, 900);
}
#[test]
fn parse_device_code_response_with_custom_interval() {
let json = r#"{
"device_code": "dc_456",
"user_code": "EFGH-5678",
"verification_uri": "https://github.com/login/device",
"expires_in": 600,
"interval": 10
}"#;
let resp: DeviceCodeResponse = serde_json::from_str(json).expect("valid JSON");
assert_eq!(resp.interval, 10);
}
// --- CopilotTokenManager ---
#[tokio::test]
async fn token_manager_caches_token_and_returns_same_value() {
// Pre-populate the cache with a token that expires far in the future.
let client = reqwest::Client::new();
let manager = CopilotTokenManager::new(client, "unused_oauth".to_string());
let far_future = unix_now() + 3600;
{
let mut guard = manager.cached.write().await;
*guard = Some(CachedCopilotToken {
token: SecretString::from("cached_session_token".to_string()),
expires_at: far_future,
});
}
let token = manager.get_token().await.expect("should return cached");
assert_eq!(token.expose_secret(), "cached_session_token");
// A second call should return the same cached token.
let token2 = manager.get_token().await.expect("should return cached");
assert_eq!(token2.expose_secret(), "cached_session_token");
}
#[tokio::test]
async fn token_manager_invalidation_clears_cache() {
let client = reqwest::Client::new();
let manager = CopilotTokenManager::new(client, "unused_oauth".to_string());
let far_future = unix_now() + 3600;
{
let mut guard = manager.cached.write().await;
*guard = Some(CachedCopilotToken {
token: SecretString::from("old_token".to_string()),
expires_at: far_future,
});
}
manager.invalidate().await;
let guard = manager.cached.read().await;
assert!(guard.is_none(), "cache should be empty after invalidation");
}
#[tokio::test]
async fn token_manager_expired_token_triggers_refresh_path() {
let client = reqwest::Client::new();
let manager = CopilotTokenManager::new(client, "unused_oauth".to_string());
// Set a token that is already expired (expires_at in the past).
{
let mut guard = manager.cached.write().await;
*guard = Some(CachedCopilotToken {
token: SecretString::from("stale_token".to_string()),
expires_at: 1, // way in the past
});
}
// get_token will try the slow path (token exchange) which will fail
// because we have no real server, but this proves the cached stale
// token is NOT returned.
let result = manager.get_token().await;
assert!(
result.is_err(),
"expired cached token should trigger exchange, which fails without a server"
);
}
#[tokio::test]
async fn token_manager_within_buffer_triggers_refresh() {
let client = reqwest::Client::new();
let manager = CopilotTokenManager::new(client, "unused_oauth".to_string());
// Set a token that expires within the refresh buffer window.
let expires_soon = unix_now() + TOKEN_REFRESH_BUFFER_SECS - 10;
{
let mut guard = manager.cached.write().await;
*guard = Some(CachedCopilotToken {
token: SecretString::from("expiring_soon".to_string()),
expires_at: expires_soon,
});
}
let result = manager.get_token().await;
assert!(
result.is_err(),
"token within buffer should trigger exchange"
);
}
// --- CopilotTokenResponse parsing ---
#[test]
fn parse_copilot_token_response() {
let json = r#"{"token": "tid=abc;exp=999;sku=123;sig=xyz", "expires_at": 1700000000}"#;
let resp: CopilotTokenResponse = serde_json::from_str(json).expect("valid JSON");
assert!(resp.token.starts_with("tid="));
assert_eq!(resp.expires_at, 1700000000);
}
}
+69
View File
@@ -18,6 +18,9 @@ pub mod config;
pub mod costs;
pub mod error;
pub mod failover;
pub mod gemini_oauth;
mod github_copilot;
pub(crate) mod github_copilot_auth;
mod nearai_chat;
pub mod oauth_helpers;
pub mod openai_codex_provider;
@@ -32,6 +35,7 @@ mod rig_adapter;
pub mod session;
pub mod smart_routing;
mod token_refreshing;
pub mod transcription;
#[cfg(test)]
mod codex_test_helpers;
@@ -48,6 +52,7 @@ pub use config::{
};
pub use error::LlmError;
pub use failover::{CooldownConfig, FailoverProvider};
pub use gemini_oauth::GeminiOauthProvider;
pub use nearai_chat::{DEFAULT_MODEL, ModelInfo, NearAiChatProvider, default_models};
pub use openai_codex_provider::OpenAiCodexProvider;
pub use openai_codex_session::{OpenAiCodexSession, OpenAiCodexSessionManager};
@@ -91,6 +96,10 @@ pub async fn create_llm_provider(
return create_llm_provider_with_config(&config.nearai, session, timeout);
}
if config.backend == "gemini_oauth" || config.backend == "gemini-oauth" {
return create_gemini_oauth_provider(config);
}
// Bedrock uses a native AWS SDK, not the rig-core registry
if config.backend == "bedrock" {
#[cfg(feature = "bedrock")]
@@ -171,6 +180,17 @@ fn create_registry_provider(
ProviderProtocol::OpenAiCompletions => create_openai_compat_from_registry(config),
ProviderProtocol::Anthropic => create_anthropic_from_registry(config),
ProviderProtocol::Ollama => create_ollama_from_registry(config),
ProviderProtocol::GithubCopilot => {
let provider =
github_copilot::GithubCopilotProvider::new(config, request_timeout_secs)?;
tracing::debug!(
provider = %config.provider_id,
model = %config.model,
base_url = %config.base_url,
"Using GitHub Copilot provider (token exchange)"
);
Ok(Arc::new(provider))
}
}
}
@@ -477,6 +497,19 @@ fn create_cheap_provider_for_backend(
});
}
if config.backend == "gemini_oauth" {
let Some(ref gemini_config) = config.gemini_oauth else {
return Err(LlmError::RequestFailed {
provider: "gemini_oauth".to_string(),
reason: "Gemini OAuth config not available for cheap model".to_string(),
});
};
let mut cheap_gemini_config = gemini_config.clone();
cheap_gemini_config.model = cheap_model.to_string();
let provider = GeminiOauthProvider::new(cheap_gemini_config)?;
return Ok(Some(Arc::new(provider)));
}
// Registry-based provider: clone config and swap model
let reg_config = config.provider.as_ref().ok_or_else(|| LlmError::RequestFailed {
provider: config.backend.clone(),
@@ -661,6 +694,17 @@ pub async fn build_provider_chain(
Ok((llm, cheap_llm, recording_handle))
}
pub fn create_gemini_oauth_provider(config: &LlmConfig) -> Result<Arc<dyn LlmProvider>, LlmError> {
let gemini_config = config
.gemini_oauth
.clone()
.ok_or_else(|| LlmError::AuthFailed {
provider: "gemini_oauth".to_string(),
})?;
let provider = gemini_oauth::GeminiOauthProvider::new(gemini_config)?;
Ok(Arc::new(provider))
}
#[cfg(test)]
mod tests {
use super::*;
@@ -692,6 +736,7 @@ mod tests {
nearai: test_nearai_config(),
provider: None,
bedrock: None,
gemini_oauth: None,
request_timeout_secs: 120,
cheap_model: None,
smart_routing_cascade: true,
@@ -773,6 +818,30 @@ mod tests {
);
}
#[test]
fn test_create_cheap_llm_provider_gemini_oauth_creates_provider() {
let mut config = test_llm_config();
config.backend = "gemini_oauth".to_string();
config.cheap_model = Some("gemini-2.5-flash-lite".to_string());
config.gemini_oauth = Some(crate::config::GeminiOauthConfig {
model: "gemini-2.5-pro".to_string(),
credentials_path: std::path::PathBuf::from("/tmp/nonexistent-creds.json"),
});
let session = Arc::new(SessionManager::new(SessionConfig::default()));
let result = create_cheap_llm_provider(&config, session);
// Should succeed and return a provider (credentials validation is deferred
// until the first LLM call, not at construction time).
let provider = result.expect("gemini_oauth cheap provider should succeed");
assert!(provider.is_some(), "Should return Some(provider)");
assert_eq!(
provider.unwrap().model_name(),
"gemini-2.5-flash-lite",
"Cheap provider should use the overridden model name"
);
}
#[test]
fn test_cheap_model_name_resolution() {
// Generic takes priority
+1
View File
@@ -344,6 +344,7 @@ pub(crate) fn build_nearai_model_fetch_config() -> crate::config::LlmConfig {
nearai: crate::config::NearAiConfig::for_model_discovery(),
provider: None,
bedrock: None,
gemini_oauth: None,
request_timeout_secs: 120,
cheap_model: None,
smart_routing_cascade: false,
+2
View File
@@ -37,6 +37,8 @@ pub enum ProviderProtocol {
Anthropic,
/// Ollama API (OpenAI-ish, no API key required).
Ollama,
/// GitHub Copilot API (OpenAI-compatible with token exchange).
GithubCopilot,
}
/// How the setup wizard should collect credentials for this provider.
+57 -7
View File
@@ -38,10 +38,49 @@ fn main() -> anyhow::Result<()> {
let _ = dotenvy::dotenv();
ironclaw::bootstrap::load_ironclaw_env();
tokio::runtime::Builder::new_multi_thread()
let result = tokio::runtime::Builder::new_multi_thread()
.enable_all()
.build()?
.block_on(async_main())
.block_on(async_main());
if let Err(ref e) = result {
format_top_level_error(e);
}
result
}
/// Format a top-level error with color and recovery hints.
fn format_top_level_error(err: &anyhow::Error) {
use ironclaw::cli::fmt;
let msg = format!("{err:#}");
eprintln!();
eprintln!(" {}\u{2717}{} {}", fmt::error(), fmt::reset(), msg);
// Provide recovery hints for common errors
let lower = msg.to_ascii_lowercase();
let hint = if lower.contains("database_url")
|| lower.contains("database") && lower.contains("not set")
{
Some("run `ironclaw onboard` or set DATABASE_URL in .env")
} else if lower.contains("connection refused") || lower.contains("connect error") {
Some("check that the database server is running")
} else if lower.contains("session") && lower.contains("not found") {
Some("run `ironclaw onboard` to set up authentication")
} else if lower.contains("secrets_master_key") {
Some("run `ironclaw onboard` or set SECRETS_MASTER_KEY in .env")
} else if lower.contains("already running") {
Some("stop the other instance or remove the stale PID file")
} else if lower.contains("onboard") {
Some("run `ironclaw onboard` to complete setup")
} else {
None
};
if let Some(hint_text) = hint {
eprintln!(" {}hint:{} {}", fmt::dim(), fmt::reset(), hint_text,);
}
eprintln!();
}
async fn async_main() -> anyhow::Result<()> {
@@ -94,6 +133,11 @@ async fn async_main() -> anyhow::Result<()> {
return ironclaw::cli::run_skills_command(skills_cmd.clone(), cli.config.as_deref())
.await;
}
Some(Command::Hooks(hooks_cmd)) => {
init_cli_tracing();
return ironclaw::cli::run_hooks_command(hooks_cmd.clone(), cli.config.as_deref())
.await;
}
Some(Command::Logs(logs_cmd)) => {
init_cli_tracing();
return ironclaw::cli::run_logs_command(logs_cmd.clone(), cli.config.as_deref()).await;
@@ -185,6 +229,7 @@ async fn async_main() -> anyhow::Result<()> {
channels_only,
provider_only,
quick,
step,
}) => {
#[cfg(any(feature = "postgres", feature = "libsql"))]
{
@@ -193,6 +238,7 @@ async fn async_main() -> anyhow::Result<()> {
channels_only: *channels_only,
provider_only: *provider_only,
quick: *quick,
steps: step.clone(),
};
let mut wizard =
SetupWizard::try_with_config_and_toml(config, cli.config.as_deref())?;
@@ -200,7 +246,7 @@ async fn async_main() -> anyhow::Result<()> {
}
#[cfg(not(any(feature = "postgres", feature = "libsql")))]
{
let _ = (skip_auth, channels_only, provider_only, quick);
let _ = (skip_auth, channels_only, provider_only, quick, step);
eprintln!("Onboarding wizard requires the 'postgres' or 'libsql' feature.");
}
return Ok(());
@@ -228,6 +274,8 @@ async fn async_main() -> anyhow::Result<()> {
}
};
let startup_start = std::time::Instant::now();
// ── Agent startup ──────────────────────────────────────────────────
// Enhanced first-run detection
@@ -686,6 +734,7 @@ async fn async_main() -> anyhow::Result<()> {
.and_then(|t| t.public_url())
.or_else(|| config.tunnel.public_url.clone()),
tunnel_provider: active_tunnel.as_ref().map(|t| t.name().to_string()),
startup_elapsed: Some(startup_start.elapsed()),
};
ironclaw::boot_screen::print_boot_screen(&boot_info);
}
@@ -797,10 +846,11 @@ async fn async_main() -> anyhow::Result<()> {
cost_guard: components.cost_guard,
sse_tx: sse_sender,
http_interceptor,
transcription: config
.transcription
.create_provider()
.map(|p| Arc::new(ironclaw::transcription::TranscriptionMiddleware::new(p))),
transcription: config.transcription.create_provider().map(|p| {
Arc::new(ironclaw::llm::transcription::TranscriptionMiddleware::new(
p,
))
}),
document_extraction: Some(Arc::new(
ironclaw::document_extraction::DocumentExtractionMiddleware::new(),
)),
+1 -1
View File
@@ -55,7 +55,7 @@ pub struct Settings {
pub secrets_master_key_hex: Option<String>,
// === Step 3: Inference Provider ===
/// LLM backend: "nearai", "anthropic", "openai", "ollama", "openai_compatible", "tinfoil", "bedrock".
/// LLM backend: "nearai", "anthropic", "openai", "github_copilot", "ollama", "openai_compatible", "tinfoil", "bedrock".
#[serde(default)]
pub llm_backend: Option<String>,
+8 -1
View File
@@ -218,6 +218,7 @@ env-var mode or skipped secrets.
| NEAR AI Cloud | API key | `llm_nearai_api_key` | `NEARAI_API_KEY` |
| Anthropic | API key | `llm_anthropic_api_key` | `ANTHROPIC_API_KEY` |
| OpenAI | API key | `llm_openai_api_key` | `OPENAI_API_KEY` |
| GitHub Copilot | OAuth token | `llm_github_copilot_token` | `GITHUB_COPILOT_TOKEN` |
| Ollama | None | - | - |
| OpenRouter | API key | `llm_openrouter_api_key` | `OPENROUTER_API_KEY` |
| OpenAI-compatible | Optional API key | `llm_compatible_api_key` | `LLM_API_KEY` |
@@ -240,6 +241,12 @@ with its own secret name and env var. It is **not** stored as `openai_compatible
5. Preserve `selected_model` on a same-backend re-run; clear it only when
switching to a different backend
**GitHub Copilot** (`setup_github_copilot`):
- Offers **GitHub device login** (recommended) or manual token paste
- Device login uses the VS Code Copilot OAuth client and stores the resulting token as `llm_github_copilot_token`
- Validates the token against `https://api.githubcopilot.com/models` before saving
- Injects `GITHUB_COPILOT_TOKEN` into the config overlay for immediate provider use
**NEAR AI** (`setup_nearai`):
- Calls `session_manager.ensure_authenticated()` which shows the auth menu:
- Options 1-2 (GitHub/Google): browser OAuth → **NEAR AI Chat** mode
@@ -530,7 +537,7 @@ pub struct Settings {
pub secrets_master_key_source: KeySource, // Keychain | Env | None
// Step 3: Inference
pub llm_backend: Option<String>, // "nearai" | "anthropic" | "openai" | "ollama" | "openai_compatible" | "bedrock"
pub llm_backend: Option<String>, // "nearai" | "anthropic" | "openai" | "github_copilot" | "ollama" | "openai_compatible" | "bedrock"
pub ollama_base_url: Option<String>,
pub openai_compatible_base_url: Option<String>,
+48 -23
View File
@@ -123,15 +123,32 @@ pub fn select_many(prompt: &str, options: &[(&str, bool)]) -> io::Result<Vec<usi
writeln!(stdout, "\r")?;
for (i, (label, _)) in options.iter().enumerate() {
let checkbox = if selected[i] { "[x]" } else { "[ ]" };
let prefix = if i == cursor_pos { ">" } else { " " };
if i == cursor_pos {
// Cursor line: cyan cursor, then colored checkbox
execute!(stdout, SetForegroundColor(Color::Cyan))?;
writeln!(stdout, " {} {} {}\r", prefix, checkbox, label)?;
write!(stdout, " \u{25b8} ")?;
if selected[i] {
execute!(stdout, SetForegroundColor(Color::Green))?;
write!(stdout, "[\u{2713}]")?;
} else {
execute!(stdout, SetForegroundColor(Color::DarkGrey))?;
write!(stdout, "[\u{00b7}]")?;
}
execute!(stdout, SetForegroundColor(Color::Cyan))?;
writeln!(stdout, " {}\r", label)?;
execute!(stdout, ResetColor)?;
} else {
writeln!(stdout, " {} {} {}\r", prefix, checkbox, label)?;
write!(stdout, " ")?;
if selected[i] {
execute!(stdout, SetForegroundColor(Color::Green))?;
write!(stdout, "[\u{2713}]")?;
execute!(stdout, ResetColor)?;
} else {
execute!(stdout, SetForegroundColor(Color::DarkGrey))?;
write!(stdout, "[\u{00b7}]")?;
execute!(stdout, ResetColor)?;
}
writeln!(stdout, " {}\r", label)?;
}
}
@@ -284,18 +301,12 @@ pub fn confirm(prompt: &str, default: bool) -> io::Result<bool> {
})
}
/// Print the IronClaw ASCII art banner in blue.
/// Print a minimal wordmark banner.
pub fn print_banner() {
let mut stdout = io::stdout();
let _ = execute!(stdout, SetForegroundColor(Color::Cyan));
use crate::cli::fmt;
println!();
println!(" {}ironclaw{}", fmt::bold_accent(), fmt::reset());
println!();
println!(r" ██╗██████╗ ██████╗ ███╗ ██╗ ██████╗██╗ █████╗ ██╗ ██╗");
println!(r" ██║██╔══██╗██╔═══██╗████╗ ██║██╔════╝██║ ██╔══██╗██║ ██║");
println!(r" ██║██████╔╝██║ ██║██╔██╗ ██║██║ ██║ ███████║██║ █╗ ██║");
println!(r" ██║██╔══██╗██║ ██║██║╚██╗██║██║ ██║ ██╔══██║██║███╗██║");
println!(r" ██║██║ ██║╚██████╔╝██║ ╚████║╚██████╗███████╗██║ ██║╚███╔███╔╝");
println!(r" ╚═╝╚═╝ ╚═╝ ╚═════╝ ╚═╝ ╚═══╝ ╚═════╝╚══════╝╚═╝ ╚═╝ ╚══╝╚══╝ ");
let _ = execute!(stdout, ResetColor);
}
/// Print a styled header box.
@@ -310,24 +321,38 @@ pub fn print_header(text: &str) {
let border = "".repeat(width);
println!();
println!("{}", border);
println!("{}", border);
println!("{}", text);
println!("{}", border);
println!("{}", border);
println!();
}
/// Print a step indicator.
/// Print a compact dot-based step indicator.
///
/// `●` = completed (green/success), `◉` = current (accent), `○` = remaining (dim).
///
/// # Example
///
/// ```ignore
/// print_step(1, 3, "NEAR AI Authentication");
/// // Output: Step 1/3: NEAR AI Authentication
/// // ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
/// print_step(3, 5, "Model Selection");
/// // Output: ● ● ◉ ○ ○ Model Selection
/// ```
pub fn print_step(current: usize, total: usize, name: &str) {
println!("Step {}/{}: {}", current, total, name);
println!("{}", "".repeat(32));
use crate::cli::fmt;
let mut dots = String::new();
for i in 1..=total {
if i > 1 {
dots.push(' ');
}
if i < current {
dots.push_str(&format!("{}\u{25CF}{}", fmt::success(), fmt::reset())); // ● green
} else if i == current {
dots.push_str(&format!("{}\u{25C9}{}", fmt::accent(), fmt::reset())); // ◉ accent
} else {
dots.push_str(&format!("{}\u{25CB}{}", fmt::dim(), fmt::reset())); // ○ dim
}
}
println!(" {} {}", dots, name);
println!();
}
+660 -227
View File
File diff suppressed because it is too large Load Diff
+233 -99
View File
@@ -56,7 +56,7 @@ use tokio::process::Command;
use crate::context::JobContext;
use crate::sandbox::{SandboxManager, SandboxPolicy};
use crate::tools::tool::{
ApprovalRequirement, Tool, ToolDomain, ToolError, ToolOutput, require_str,
ApprovalRequirement, RiskLevel, Tool, ToolDomain, ToolError, ToolOutput, require_str,
};
/// Maximum output size before truncation (64KB).
@@ -117,7 +117,7 @@ static NEVER_AUTO_APPROVE_PATTERNS: LazyLock<Vec<&'static str>> = LazyLock::new(
"init 0",
"init 6",
"iptables",
"nft ",
"nft",
"useradd",
"userdel",
"passwd",
@@ -132,6 +132,7 @@ static NEVER_AUTO_APPROVE_PATTERNS: LazyLock<Vec<&'static str>> = LazyLock::new(
"docker rmi",
"docker system prune",
"git push --force",
"git push --force-with-lease",
"git push -f",
"git reset --hard",
"git clean -f",
@@ -139,6 +140,7 @@ static NEVER_AUTO_APPROVE_PATTERNS: LazyLock<Vec<&'static str>> = LazyLock::new(
"DROP DATABASE",
"TRUNCATE",
"DELETE FROM",
"sudo",
]
});
@@ -195,15 +197,205 @@ const SAFE_ENV_VARS: &[&str] = &[
"WINDIR",
];
/// Check whether a shell command contains patterns that must never be auto-approved.
/// Low-risk command prefixes: strictly read-only commands with no side effects.
/// Note: `sed`, `awk`, and `find` are intentionally excluded — they have destructive
/// modes (`sed -i`, `awk -i inplace`, `find -delete`) and are classified as Medium.
static LOW_RISK_PATTERNS: LazyLock<Vec<&'static str>> = LazyLock::new(|| {
vec![
"ls",
"ll",
"la",
"dir",
"cat",
"less",
"more",
"head",
"tail",
"grep",
"rg",
"ag",
"fd",
"locate",
"echo",
"printf",
"pwd",
"cd",
"env",
"printenv",
"which",
"whereis",
"type",
"date",
"cal",
"uptime",
"uname",
"df",
"du",
"free",
"top",
"htop",
"ps",
"git status",
"git log",
"git diff",
"git show",
"git branch",
"git remote",
"git fetch",
"cargo check",
"cargo clippy",
"curl --head",
"curl -I",
"ping",
"wc",
"sort",
"uniq",
"tr",
"cut",
"jq",
"yq",
"file",
"stat",
"man",
]
});
/// Medium-risk command prefixes: mutations that are generally reversible, plus commands with
/// potentially destructive flags (e.g. `sed -i`, `awk -i inplace`, `find -delete`).
static MEDIUM_RISK_PATTERNS: LazyLock<Vec<&'static str>> = LazyLock::new(|| {
vec![
// Text processors with in-place/destructive modes
"awk",
"sed",
"find",
"mkdir",
"rmdir",
"touch",
"cp",
"copy",
"mv",
"move",
"git commit",
"git add",
"git push",
"git checkout",
"git switch",
"git merge",
"git rebase",
"git stash",
"git tag",
"cargo build",
"cargo run",
"cargo test",
"npm test",
"npm run test",
"yarn test",
"npm install",
"npm ci",
"npm update",
"pip install",
"pip uninstall",
"brew install",
"brew uninstall",
"apt install",
"apt remove",
"make",
"cmake",
"tar",
"zip",
"unzip",
"gzip",
"gunzip",
"ssh",
"scp",
"rsync",
"curl",
"wget",
"docker build",
"docker pull",
"docker run",
"kubectl apply",
"kubectl create",
]
});
/// Match a pipeline segment against a risk pattern using word-boundary rules.
///
/// Even when the user has chosen "always approve" for the shell tool, these commands
/// require explicit per-invocation approval because they are destructive.
pub fn requires_explicit_approval(command: &str) -> bool {
let lower = command.to_lowercase();
NEVER_AUTO_APPROVE_PATTERNS
.iter()
.any(|p| lower.contains(&p.to_lowercase()))
/// - **Multi-word patterns** (e.g. `"git status"`): the segment must equal the
/// pattern or start with `"<pattern> "`, so `"git statusbar"` does not match
/// `"git status"`.
/// - **Single-word patterns** (e.g. `"ls"`): the first whitespace-delimited
/// token of the segment must equal the pattern exactly, so `"lsblk"` does
/// not match `"ls"`.
fn matches_command_pattern(segment: &str, pattern: &str) -> bool {
if pattern.contains(' ') {
segment == pattern || segment.starts_with(&format!("{} ", pattern))
} else {
segment.split_whitespace().next().unwrap_or("") == pattern
}
}
/// Classify a shell command into a [`RiskLevel`].
///
/// The command is split on `|`, `&`, `;` and each segment is classified
/// independently; the overall risk is the **maximum** across all segments
/// so a dangerous sub-command in a pipeline is never missed.
///
/// Per-segment priority (highest wins):
/// 1. **High** — segment matches [`NEVER_AUTO_APPROVE_PATTERNS`] (destructive / irreversible).
/// 2. **Low** — segment matches [`LOW_RISK_PATTERNS`] (strictly read-only).
/// 3. **Medium** — segment matches [`MEDIUM_RISK_PATTERNS`] (reversible mutations).
/// 4. **Medium** — unknown commands default to Medium (safer than auto-approving).
///
/// All matching uses word-boundary rules (see [`matches_command_pattern`]) to
/// prevent false positives like `"makeshutdownscript"` matching `"shutdown"` or
/// `"lsblk"` matching `"ls"`.
pub fn classify_command_risk(command: &str) -> RiskLevel {
// For pipelines/chains, take the maximum risk across all segments.
command
.split(['|', '&', ';'])
.map(str::trim)
.filter(|s| !s.is_empty())
.map(|segment| {
let seg_lower = segment.to_lowercase();
if NEVER_AUTO_APPROVE_PATTERNS
.iter()
.any(|p| matches_command_pattern(&seg_lower, &p.to_lowercase()))
{
RiskLevel::High
} else if LOW_RISK_PATTERNS
.iter()
.any(|p| matches_command_pattern(&seg_lower, p))
{
RiskLevel::Low
} else if MEDIUM_RISK_PATTERNS
.iter()
.any(|p| matches_command_pattern(&seg_lower, p))
{
RiskLevel::Medium
} else {
// Unknown commands default to Medium (safer than auto-approving).
RiskLevel::Medium
}
})
.max()
.unwrap_or(RiskLevel::Medium)
}
/// Extract the `command` field from a tool-call parameter value.
///
/// Handles both the normal case (a JSON object with a `"command"` key) and the
/// rare case where the LLM provider returns string-encoded JSON.
fn extract_command_param(params: &serde_json::Value) -> Option<String> {
params
.get("command")
.and_then(|c| c.as_str().map(String::from))
.or_else(|| {
params
.as_str()
.and_then(|s| serde_json::from_str::<serde_json::Value>(s).ok())
.and_then(|v| v.get("command").and_then(|c| c.as_str().map(String::from)))
})
}
/// Detect command injection and obfuscation attempts.
@@ -698,24 +890,24 @@ impl Tool for ShellTool {
Ok(ToolOutput::success(result, duration))
}
fn risk_level_for(&self, params: &serde_json::Value) -> RiskLevel {
extract_command_param(params)
.map(|cmd| classify_command_risk(&cmd))
.unwrap_or(RiskLevel::Medium)
}
fn requires_approval(&self, params: &serde_json::Value) -> ApprovalRequirement {
let cmd = params
.get("command")
.and_then(|c| c.as_str().map(String::from))
.or_else(|| {
params
.as_str()
.and_then(|s| serde_json::from_str::<serde_json::Value>(s).ok())
.and_then(|v| v.get("command").and_then(|c| c.as_str().map(String::from)))
});
if let Some(ref cmd) = cmd
&& requires_explicit_approval(cmd)
{
return ApprovalRequirement::Always;
match self.risk_level_for(params) {
// Low maps to UnlessAutoApproved rather than Never: shell redirections
// (e.g. `cat /etc/shadow > /tmp/out`) are not split on `>`, so a Low command
// with a redirect would bypass approval entirely with Never. Keeping
// UnlessAutoApproved preserves the graduated metadata for audit while
// ensuring approval policy stays conservative until redirect-aware parsing
// is in place.
RiskLevel::Low => ApprovalRequirement::UnlessAutoApproved,
RiskLevel::Medium => ApprovalRequirement::UnlessAutoApproved,
RiskLevel::High => ApprovalRequirement::Always,
}
ApprovalRequirement::UnlessAutoApproved
}
fn requires_sanitization(&self) -> bool {
@@ -799,74 +991,11 @@ mod tests {
assert!(matches!(result, Err(ToolError::Timeout(_))));
}
#[test]
fn test_requires_explicit_approval() {
// Destructive commands should require explicit approval
assert!(requires_explicit_approval("rm -rf /tmp/stuff"));
assert!(requires_explicit_approval("git push --force origin main"));
assert!(requires_explicit_approval("git reset --hard HEAD~5"));
assert!(requires_explicit_approval("docker rm container_name"));
assert!(requires_explicit_approval("kill -9 12345"));
assert!(requires_explicit_approval("DROP TABLE users;"));
// Safe commands should not
assert!(!requires_explicit_approval("cargo build"));
assert!(!requires_explicit_approval("git status"));
assert!(!requires_explicit_approval("ls -la"));
assert!(!requires_explicit_approval("echo hello"));
assert!(!requires_explicit_approval("cat file.txt"));
assert!(!requires_explicit_approval(
"git push origin feature-branch"
));
}
/// Replicate the extraction logic from agent_loop.rs to prove it works
/// when `arguments` is a `serde_json::Value::Object` (the common case
/// that was previously broken because `Value::Object.as_str()` returns None).
#[test]
fn test_destructive_command_extraction_from_object_args() {
let arguments = serde_json::json!({"command": "rm -rf /tmp/stuff"});
let cmd = arguments
.get("command")
.and_then(|c| c.as_str().map(String::from))
.or_else(|| {
arguments
.as_str()
.and_then(|s| serde_json::from_str::<serde_json::Value>(s).ok())
.and_then(|v| v.get("command").and_then(|c| c.as_str().map(String::from)))
});
assert_eq!(cmd.as_deref(), Some("rm -rf /tmp/stuff"));
assert!(requires_explicit_approval(cmd.as_deref().unwrap()));
}
/// Verify extraction still works when `arguments` is a JSON string
/// (rare, but possible if the LLM provider returns string-encoded JSON).
#[test]
fn test_destructive_command_extraction_from_string_args() {
let arguments =
serde_json::Value::String(r#"{"command": "git push --force origin main"}"#.to_string());
let cmd = arguments
.get("command")
.and_then(|c| c.as_str().map(String::from))
.or_else(|| {
arguments
.as_str()
.and_then(|s| serde_json::from_str::<serde_json::Value>(s).ok())
.and_then(|v| v.get("command").and_then(|c| c.as_str().map(String::from)))
});
assert_eq!(cmd.as_deref(), Some("git push --force origin main"));
assert!(requires_explicit_approval(cmd.as_deref().unwrap()));
}
#[test]
fn test_requires_approval_destructive_command() {
use crate::tools::tool::ApprovalRequirement;
let tool = ShellTool::new();
// Destructive commands must return Always to bypass auto-approve.
// High-risk commands must return Always to bypass auto-approve.
assert_eq!(
tool.requires_approval(&serde_json::json!({"command": "rm -rf /tmp"})),
ApprovalRequirement::Always
@@ -885,15 +1014,17 @@ mod tests {
fn test_requires_approval_safe_command() {
use crate::tools::tool::ApprovalRequirement;
let tool = ShellTool::new();
// Safe commands return UnlessAutoApproved (can be auto-approved).
// Medium-risk commands return UnlessAutoApproved (can be auto-approved).
assert_eq!(
tool.requires_approval(&serde_json::json!({"command": "cargo build"})),
ApprovalRequirement::UnlessAutoApproved
);
assert_eq!(
tool.requires_approval(&serde_json::json!({"command": "echo hello"})),
ApprovalRequirement::UnlessAutoApproved
);
// Low-risk commands also return UnlessAutoApproved (conservative until
// redirect-aware parsing is in place — see RiskLevel::Low mapping comment).
let r_echo = tool.requires_approval(&serde_json::json!({"command": "echo hello"}));
assert_eq!(r_echo, ApprovalRequirement::UnlessAutoApproved); // safety: test code
let r_ls = tool.requires_approval(&serde_json::json!({"command": "ls -la"}));
assert_eq!(r_ls, ApprovalRequirement::UnlessAutoApproved); // safety: test code
}
#[test]
@@ -1370,9 +1501,12 @@ mod tests {
#[test]
fn test_approval_with_mixed_case_destructive() {
// Case-insensitive destructive command detection
assert!(requires_explicit_approval("RM -RF /tmp"));
assert!(requires_explicit_approval("Git Push --Force origin main"));
assert!(requires_explicit_approval("DROP table users;"));
// Case-insensitive destructive command detection → must be High risk
let r1 = classify_command_risk("RM -RF /tmp");
assert_eq!(r1, RiskLevel::High); // safety: test code
let r2 = classify_command_risk("Git Push --Force origin main");
assert_eq!(r2, RiskLevel::High); // safety: test code
let r3 = classify_command_risk("DROP table users;");
assert_eq!(r3, RiskLevel::High); // safety: test code
}
}
+17 -5
View File
@@ -45,11 +45,23 @@ impl ToolInfoDetail {
}
fn schema_param_names(schema: &serde_json::Value) -> Vec<String> {
schema
.get("properties")
.and_then(|p| p.as_object())
.map(|props| props.keys().cloned().collect())
.unwrap_or_default()
let mut names = std::collections::BTreeSet::new();
if let Some(props) = schema.get("properties").and_then(|p| p.as_object()) {
names.extend(props.keys().cloned());
}
for key in ["allOf", "oneOf", "anyOf"] {
if let Some(variants) = schema.get(key).and_then(|v| v.as_array()) {
for variant in variants {
if let Some(props) = variant.get("properties").and_then(|p| p.as_object()) {
names.extend(props.keys().cloned());
}
}
}
}
names.into_iter().collect()
}
fn fallback_summary(schema: &serde_json::Value) -> ToolDiscoverySummary {
+701 -12
View File
@@ -1,4 +1,4 @@
pub(crate) fn prepare_tool_params(
pub fn prepare_tool_params(
tool: &dyn crate::tools::tool::Tool,
params: &serde_json::Value,
) -> serde_json::Value {
@@ -9,14 +9,87 @@ pub(crate) fn prepare_params_for_schema(
params: &serde_json::Value,
schema: &serde_json::Value,
) -> serde_json::Value {
coerce_value(params, schema)
let resolved = resolve_refs(schema);
coerce_value(params, &resolved)
}
// ── $ref resolution ──────────────────────────────────────────────────
/// Inline all `$ref` pointers in a JSON Schema so downstream coercion
/// operates on a flat, self-contained schema tree.
///
/// Supports `#/definitions/<name>` and `#/$defs/<name>` (JSON Schema
/// draft-07 and 2020-12 respectively). Unknown `$ref` formats are left
/// unchanged. A depth limit prevents infinite recursion from circular refs.
fn resolve_refs(schema: &serde_json::Value) -> serde_json::Value {
let definitions = schema
.get("definitions")
.or_else(|| schema.get("$defs"))
.cloned()
.unwrap_or(serde_json::Value::Null);
resolve_refs_inner(schema, &definitions, 0)
}
const MAX_REF_DEPTH: usize = 16;
fn resolve_refs_inner(
schema: &serde_json::Value,
definitions: &serde_json::Value,
depth: usize,
) -> serde_json::Value {
if depth > MAX_REF_DEPTH {
return schema.clone();
}
match schema {
serde_json::Value::Object(obj) => {
// If this node is a $ref, resolve it and recurse into the target.
if let Some(ref_str) = obj.get("$ref").and_then(|v| v.as_str()) {
if let Some(target) = resolve_ref_pointer(ref_str, definitions) {
return resolve_refs_inner(&target, definitions, depth + 1);
}
return schema.clone();
}
// Recursively resolve refs in all values (skip definitions maps).
let resolved: serde_json::Map<String, serde_json::Value> = obj
.iter()
.map(|(k, v)| {
if k == "definitions" || k == "$defs" {
(k.clone(), v.clone())
} else {
(k.clone(), resolve_refs_inner(v, definitions, depth + 1))
}
})
.collect();
serde_json::Value::Object(resolved)
}
serde_json::Value::Array(arr) => serde_json::Value::Array(
arr.iter()
.map(|v| resolve_refs_inner(v, definitions, depth + 1))
.collect(),
),
_ => schema.clone(),
}
}
fn resolve_ref_pointer(
ref_str: &str,
definitions: &serde_json::Value,
) -> Option<serde_json::Value> {
let path = ref_str.strip_prefix("#/")?;
let parts: Vec<&str> = path.split('/').collect();
if parts.len() == 2 && (parts[0] == "definitions" || parts[0] == "$defs") {
return definitions.get(parts[1]).cloned();
}
None
}
// ── Core coercion ────────────────────────────────────────────────────
fn coerce_value(value: &serde_json::Value, schema: &serde_json::Value) -> serde_json::Value {
// This coercer intentionally handles the concrete schema shapes we expose in
// discovery today. It does not resolve combinators like anyOf/oneOf/allOf or
// references via $ref; those schemas pass through unchanged unless they also
// advertise a directly coercible type/property shape.
// This coercer handles concrete schema shapes including discriminated unions
// (oneOf/anyOf with const or single-element enum discriminators), allOf
// merges, and $ref references (resolved in a pre-pass).
if value.is_null() {
return value.clone();
}
@@ -47,12 +120,35 @@ fn coerce_value(value: &serde_json::Value, schema: &serde_json::Value) -> serde_
return value.clone();
}
let properties = schema.get("properties").and_then(|p| p.as_object());
let additional_schema = schema.get("additionalProperties").filter(|v| v.is_object());
let resolved = resolve_effective_properties(schema, obj);
let properties = resolved
.as_ref()
.or_else(|| schema.get("properties").and_then(|p| p.as_object()));
let additional_schema = schema
.get("additionalProperties")
.filter(|v| v.is_object())
.or_else(|| resolve_additional_properties(schema, obj));
let required: std::collections::HashSet<&str> = schema
.get("required")
.and_then(|r| r.as_array())
.map(|arr| arr.iter().filter_map(|v| v.as_str()).collect())
.unwrap_or_default();
let mut coerced = obj.clone();
for (key, current) in &mut coerced {
if let Some(prop_schema) = properties.and_then(|props| props.get(key)) {
// LLMs send "" for optional fields instead of omitting them.
// Coerce to null only when the field is not required AND the schema
// allows null or doesn't allow string — a `type: "string"` field
// may legitimately accept "" as a meaningful value.
if current.as_str() == Some("")
&& !required.contains(key.as_str())
&& (schema_allows_type(prop_schema, "null")
|| !schema_allows_type(prop_schema, "string"))
{
*current = serde_json::Value::Null;
continue;
}
*current = coerce_value(current, prop_schema);
continue;
}
@@ -68,11 +164,179 @@ fn coerce_value(value: &serde_json::Value, schema: &serde_json::Value) -> serde_
value.clone()
}
/// When the schema uses `oneOf`, `anyOf`, or `allOf` combinators, build a
/// merged property map that can be used for coercion.
///
/// - Top-level `properties` are included first (base properties).
/// - `allOf`: merge ALL variants' properties (last-wins on conflicts).
/// - `oneOf`/`anyOf`: find the discriminated match and merge its properties.
///
/// Returns `None` if no combinators are present or no match is found, so the
/// caller falls back to the existing top-level `properties` lookup.
fn resolve_effective_properties(
schema: &serde_json::Value,
obj: &serde_json::Map<String, serde_json::Value>,
) -> Option<serde_json::Map<String, serde_json::Value>> {
collect_properties(schema, obj, 0)
}
const MAX_COMBINATOR_DEPTH: usize = 4;
/// Recursively collect properties from a schema and its combinator variants.
fn collect_properties(
schema: &serde_json::Value,
obj: &serde_json::Map<String, serde_json::Value>,
depth: usize,
) -> Option<serde_json::Map<String, serde_json::Value>> {
if depth > MAX_COMBINATOR_DEPTH {
return None;
}
let has_combinators = schema.get("allOf").is_some()
|| schema.get("oneOf").is_some()
|| schema.get("anyOf").is_some();
if !has_combinators {
return None;
}
let mut merged = serde_json::Map::new();
// Start with top-level properties
if let Some(props) = schema.get("properties").and_then(|p| p.as_object()) {
merged.extend(props.iter().map(|(k, v)| (k.clone(), v.clone())));
}
// allOf: merge ALL variants' properties, recursing into nested combinators
if let Some(all_of) = schema.get("allOf").and_then(|a| a.as_array()) {
for variant in all_of {
if let Some(props) = variant.get("properties").and_then(|p| p.as_object()) {
merged.extend(props.iter().map(|(k, v)| (k.clone(), v.clone())));
}
// Recurse into variant if it has its own combinators
if let Some(nested) = collect_properties(variant, obj, depth + 1) {
merged.extend(nested);
}
}
}
// oneOf/anyOf: find discriminated match and merge its properties
for key in ["oneOf", "anyOf"] {
if let Some(variants) = schema.get(key).and_then(|v| v.as_array())
&& let Some(variant) = find_discriminated_variant(variants, obj)
{
if let Some(props) = variant.get("properties").and_then(|p| p.as_object()) {
merged.extend(props.iter().map(|(k, v)| (k.clone(), v.clone())));
}
// Recurse into matched variant if it has its own combinators
if let Some(nested) = collect_properties(variant, obj, depth + 1) {
merged.extend(nested);
}
}
}
if merged.is_empty() {
None
} else {
Some(merged)
}
}
/// Find `additionalProperties` from a matched combinator variant.
///
/// Checks `allOf` variants first (last-wins), then the matched `oneOf`/`anyOf`
/// variant. Returns `None` if no variant defines `additionalProperties`.
fn resolve_additional_properties<'a>(
schema: &'a serde_json::Value,
obj: &serde_json::Map<String, serde_json::Value>,
) -> Option<&'a serde_json::Value> {
// allOf: last variant with additionalProperties wins
if let Some(all_of) = schema.get("allOf").and_then(|a| a.as_array()) {
for variant in all_of.iter().rev() {
if let Some(ap) = variant.get("additionalProperties")
&& ap.is_object()
{
return Some(ap);
}
}
}
// oneOf/anyOf: check matched variant
for key in ["oneOf", "anyOf"] {
if let Some(variants) = schema.get(key).and_then(|v| v.as_array())
&& let Some(variant) = find_discriminated_variant(variants, obj)
&& let Some(ap) = variant.get("additionalProperties")
&& ap.is_object()
{
return Some(ap);
}
}
None
}
/// Find a `oneOf`/`anyOf` variant that matches the given object by checking
/// `const`-valued and single-element `enum`-valued properties (discriminators).
///
/// A variant matches when ALL its discriminator properties match the object's
/// values and at least one such discriminator exists. Returns `None` if no
/// variant matches (safe fallback — no coercion).
fn find_discriminated_variant<'a>(
variants: &'a [serde_json::Value],
obj: &serde_json::Map<String, serde_json::Value>,
) -> Option<&'a serde_json::Value> {
variants.iter().find(|variant| {
let Some(props) = variant.get("properties").and_then(|p| p.as_object()) else {
return false;
};
let mut discriminator_count = 0;
for (key, prop_schema) in props {
// Check for const discriminator
if let Some(const_val) = prop_schema.get("const") {
discriminator_count += 1;
match obj.get(key) {
Some(v) if v == const_val => {}
_ => return false,
}
continue;
}
// Check for single-element enum discriminator
if let Some(enum_vals) = prop_schema.get("enum").and_then(|e| e.as_array())
&& enum_vals.len() == 1
{
discriminator_count += 1;
match obj.get(key) {
Some(v) if v == &enum_vals[0] => {}
_ => return false,
}
}
}
discriminator_count > 0
})
}
fn coerce_string_value(s: &str, schema: &serde_json::Value) -> Option<serde_json::Value> {
// LLMs often send "" instead of null for optional fields. Coerce empty
// strings to null when the schema allows null but not string, or allows
// both but the value is empty (a string field with content "" is kept).
if s.is_empty() && schema_allows_type(schema, "null") && !schema_allows_type(schema, "string") {
return Some(serde_json::Value::Null);
}
if schema_allows_type(schema, "string") {
return None;
}
// Empty string with no type match — return unchanged since we can't
// determine the intended type.
if s.is_empty() {
return None;
}
if schema_allows_type(schema, "integer")
&& let Ok(v) = s.parse::<i64>()
{
@@ -114,10 +378,15 @@ fn schema_allows_type(schema: &serde_json::Value, expected: &str) -> bool {
Some(serde_json::Value::String(t)) => t == expected,
Some(serde_json::Value::Array(types)) => types.iter().any(|t| t.as_str() == Some(expected)),
_ => match expected {
"object" => schema
.get("properties")
.and_then(|p| p.as_object())
.is_some(),
"object" => {
schema
.get("properties")
.and_then(|p| p.as_object())
.is_some()
|| schema.get("oneOf").is_some()
|| schema.get("anyOf").is_some()
|| schema.get("allOf").is_some()
}
"array" => schema.get("items").is_some(),
_ => false,
},
@@ -325,6 +594,91 @@ mod tests {
assert_eq!(result["value"], serde_json::json!("{\"mode\":\"raw\"}")); // safety: test-only assertion
}
#[test]
fn coerces_empty_string_to_null_for_nullable_non_required_field() {
let schema = serde_json::json!({
"type": "object",
"properties": {
"timezone": { "type": ["string", "null"] },
"schedule": { "type": "string" }
},
"required": ["schedule"]
});
let params = serde_json::json!({
"timezone": "",
"schedule": "0 9 * * *"
});
let result = prepare_params_for_schema(&params, &schema);
// Non-required nullable "timezone" with empty string → null
assert_eq!(result["timezone"], serde_json::Value::Null);
// Required "schedule" keeps its value even if empty would be weird
assert_eq!(result["schedule"], serde_json::json!("0 9 * * *"));
}
#[test]
fn keeps_empty_string_for_non_required_string_only_field() {
let schema = serde_json::json!({
"type": "object",
"properties": {
"timezone": { "type": "string" },
"schedule": { "type": "string" }
},
"required": ["schedule"]
});
let params = serde_json::json!({
"timezone": "",
"schedule": "0 9 * * *"
});
let result = prepare_params_for_schema(&params, &schema);
// Non-required string-only "timezone" keeps empty string (meaningful value)
assert_eq!(result["timezone"], serde_json::json!(""));
assert_eq!(result["schedule"], serde_json::json!("0 9 * * *"));
}
#[test]
fn coerces_empty_string_to_null_for_explicit_nullable_type() {
let schema = serde_json::json!({
"type": "object",
"properties": {
"from_timezone": { "type": ["string", "null"] },
"operation": { "type": "string" }
},
"required": ["operation"]
});
let params = serde_json::json!({
"from_timezone": "",
"operation": "now"
});
let result = prepare_params_for_schema(&params, &schema);
// Nullable type with empty string → null (even if it were required,
// the per-value coercion in coerce_string_value handles this)
assert_eq!(result["from_timezone"], serde_json::Value::Null);
assert_eq!(result["operation"], serde_json::json!("now"));
}
#[test]
fn keeps_empty_string_for_required_string_only_field() {
let schema = serde_json::json!({
"type": "object",
"properties": {
"name": { "type": "string" }
},
"required": ["name"]
});
let params = serde_json::json!({ "name": "" });
let result = prepare_params_for_schema(&params, &schema);
// Required string-only field keeps empty string
assert_eq!(result["name"], serde_json::json!(""));
}
#[test]
fn permissive_schema_is_noop() {
let schema = serde_json::json!({
@@ -339,6 +693,341 @@ mod tests {
assert_eq!(result["count"], serde_json::json!("10")); // safety: test-only assertion
}
#[test]
fn coerces_oneof_discriminated_variant() {
let schema = serde_json::json!({
"oneOf": [
{
"type": "object",
"properties": {
"action": { "const": "list_repos" },
"limit": { "type": "integer" },
"sort": { "type": "string" }
}
},
{
"type": "object",
"properties": {
"action": { "const": "get_repo" },
"repo": { "type": "string" }
}
}
]
});
let params = serde_json::json!({
"action": "list_repos",
"limit": "100",
"sort": "stars"
});
let result = prepare_params_for_schema(&params, &schema);
assert_eq!(result["action"], serde_json::json!("list_repos"));
assert_eq!(result["limit"], serde_json::json!(100));
assert_eq!(result["sort"], serde_json::json!("stars"));
}
#[test]
fn coerces_oneof_with_enum_discriminator() {
let schema = serde_json::json!({
"oneOf": [
{
"type": "object",
"properties": {
"mode": { "enum": ["fetch"] },
"count": { "type": "integer" }
}
},
{
"type": "object",
"properties": {
"mode": { "enum": ["push"] },
"force": { "type": "boolean" }
}
}
]
});
let params = serde_json::json!({
"mode": "push",
"force": "true"
});
let result = prepare_params_for_schema(&params, &schema);
assert_eq!(result["mode"], serde_json::json!("push"));
assert_eq!(result["force"], serde_json::json!(true));
}
#[test]
fn coerces_allof_merged_properties() {
let schema = serde_json::json!({
"allOf": [
{
"type": "object",
"properties": {
"page": { "type": "integer" }
}
},
{
"type": "object",
"properties": {
"per_page": { "type": "integer" },
"verbose": { "type": "boolean" }
}
}
]
});
let params = serde_json::json!({
"page": "2",
"per_page": "50",
"verbose": "false"
});
let result = prepare_params_for_schema(&params, &schema);
assert_eq!(result["page"], serde_json::json!(2));
assert_eq!(result["per_page"], serde_json::json!(50));
assert_eq!(result["verbose"], serde_json::json!(false));
}
#[test]
fn oneof_no_discriminator_match_is_noop() {
let schema = serde_json::json!({
"oneOf": [
{
"type": "object",
"properties": {
"action": { "const": "list_repos" },
"limit": { "type": "integer" }
}
},
{
"type": "object",
"properties": {
"action": { "const": "get_repo" },
"repo": { "type": "string" }
}
}
]
});
let params = serde_json::json!({
"action": "unknown_action",
"limit": "100"
});
let result = prepare_params_for_schema(&params, &schema);
// No variant matched, so no coercion happens
assert_eq!(result["limit"], serde_json::json!("100"));
}
#[test]
fn anyof_without_discriminator_is_noop() {
let schema = serde_json::json!({
"anyOf": [
{
"type": "object",
"properties": {
"name": { "type": "string" }
},
"required": ["name"]
},
{
"type": "object",
"properties": {
"id": { "type": "integer" }
},
"required": ["id"]
}
]
});
let params = serde_json::json!({
"id": "42"
});
let result = prepare_params_for_schema(&params, &schema);
// No const/enum discriminators, so no variant matches, no coercion
assert_eq!(result["id"], serde_json::json!("42"));
}
#[test]
fn resolves_ref_and_coerces_referenced_properties() {
let schema = serde_json::json!({
"type": "object",
"definitions": {
"Pagination": {
"type": "object",
"properties": {
"page": { "type": "integer" },
"per_page": { "type": "integer" }
}
}
},
"allOf": [
{ "$ref": "#/definitions/Pagination" },
{
"type": "object",
"properties": {
"query": { "type": "string" }
}
}
]
});
let params = serde_json::json!({
"page": "2",
"per_page": "50",
"query": "test"
});
let result = prepare_params_for_schema(&params, &schema);
assert_eq!(result["page"], serde_json::json!(2));
assert_eq!(result["per_page"], serde_json::json!(50));
assert_eq!(result["query"], serde_json::json!("test"));
}
#[test]
fn resolves_nested_refs_in_oneof_variants() {
let schema = serde_json::json!({
"type": "object",
"$defs": {
"ListParams": {
"properties": {
"action": { "const": "list" },
"limit": { "type": "integer" }
}
}
},
"oneOf": [
{ "$ref": "#/$defs/ListParams" },
{
"properties": {
"action": { "const": "get" },
"id": { "type": "integer" }
}
}
]
});
let params = serde_json::json!({
"action": "list",
"limit": "25"
});
let result = prepare_params_for_schema(&params, &schema);
assert_eq!(result["limit"], serde_json::json!(25));
}
#[test]
fn coerces_nested_combinators_allof_containing_oneof() {
// allOf where one variant is itself a oneOf (nested combinator)
let schema = serde_json::json!({
"type": "object",
"allOf": [
{
"properties": {
"version": { "type": "integer" }
}
},
{
"oneOf": [
{
"properties": {
"mode": { "const": "fast" },
"threads": { "type": "integer" }
}
},
{
"properties": {
"mode": { "const": "safe" },
"retries": { "type": "integer" }
}
}
]
}
]
});
let params = serde_json::json!({
"version": "3",
"mode": "fast",
"threads": "8"
});
let result = prepare_params_for_schema(&params, &schema);
assert_eq!(result["version"], serde_json::json!(3));
assert_eq!(result["threads"], serde_json::json!(8));
}
#[test]
fn coerces_array_items_with_oneof_discriminator() {
let schema = serde_json::json!({
"type": "object",
"properties": {
"actions": {
"type": "array",
"items": {
"oneOf": [
{
"type": "object",
"properties": {
"type": { "const": "move" },
"distance": { "type": "integer" }
}
},
{
"type": "object",
"properties": {
"type": { "const": "wait" },
"seconds": { "type": "number" }
}
}
]
}
}
}
});
let params = serde_json::json!({
"actions": [
{ "type": "move", "distance": "10" },
{ "type": "wait", "seconds": "2.5" }
]
});
let result = prepare_params_for_schema(&params, &schema);
assert_eq!(result["actions"][0]["distance"], serde_json::json!(10));
assert_eq!(result["actions"][1]["seconds"], serde_json::json!(2.5));
}
#[test]
fn circular_ref_does_not_infinite_loop() {
let schema = serde_json::json!({
"type": "object",
"definitions": {
"Node": {
"type": "object",
"properties": {
"value": { "type": "integer" },
"child": { "$ref": "#/definitions/Node" }
}
}
},
"properties": {
"root": { "$ref": "#/definitions/Node" }
}
});
let params = serde_json::json!({
"root": { "value": "42" }
});
// Should not hang — depth limit stops the recursion
let result = prepare_params_for_schema(&params, &schema);
assert_eq!(result["root"]["value"], serde_json::json!(42));
}
#[test]
fn prepare_tool_params_uses_discovery_schema() {
let tool = StubTool {
+1 -1
View File
@@ -133,7 +133,7 @@ pub fn process_tool_result(
let content = match result {
Ok(output) => {
let sanitized = safety.sanitize_tool_output(tool_name, output);
safety.wrap_for_llm(tool_name, &sanitized.content, sanitized.was_modified)
safety.wrap_for_llm(tool_name, &sanitized.content)
}
Err(e) => format!("Error: {}", e),
};
+1 -1
View File
@@ -34,6 +34,6 @@ pub(crate) use coercion::prepare_tool_params;
pub use rate_limiter::RateLimiter;
pub use registry::ToolRegistry;
pub use tool::{
ApprovalContext, ApprovalRequirement, Tool, ToolDomain, ToolError, ToolOutput,
ApprovalContext, ApprovalRequirement, RiskLevel, Tool, ToolDomain, ToolError, ToolOutput,
ToolRateLimitConfig, redact_params, validate_tool_schema,
};
+1 -1
View File
@@ -604,7 +604,7 @@ impl ToolRegistry {
self.register(Arc::new(BuildSoftwareTool::new(Arc::clone(&builder))))
.await;
tracing::info!("Registered software builder tool");
tracing::debug!("Registered software builder tool");
builder
}
+83 -5
View File
@@ -42,11 +42,38 @@ pub fn validate_strict_schema(
}
}
/// Returns true if the schema uses `oneOf`, `anyOf`, or `allOf` combinators
/// where at least one variant is an object type (has `type: "object"` or `properties`).
fn has_object_combinator_variants(schema: &serde_json::Value) -> bool {
for key in ["oneOf", "anyOf", "allOf"] {
if let Some(variants) = schema.get(key).and_then(|v| v.as_array())
&& variants.iter().any(|v| {
v.get("type").and_then(|t| t.as_str()) == Some("object")
|| v.get("properties").is_some()
})
{
return true;
}
}
false
}
/// Recursively validate an object-typed schema node.
fn check_object_schema(schema: &serde_json::Value, path: &str) -> Vec<String> {
let mut errors = Vec::new();
// Rule 1: must have "type": "object"
// Report non-array combinator values as errors.
for key in ["oneOf", "anyOf", "allOf"] {
if let Some(val) = schema.get(key)
&& !val.is_array()
{
errors.push(format!("{path}: \"{key}\" must be an array"));
}
}
let has_combinators = has_object_combinator_variants(schema);
// Rule 1: must have "type": "object" (unless combinators define the structure)
match schema.get("type").and_then(|t| t.as_str()) {
Some("object") => {}
Some(other) => {
@@ -54,16 +81,67 @@ fn check_object_schema(schema: &serde_json::Value, path: &str) -> Vec<String> {
return errors;
}
None => {
errors.push(format!("{path}: missing \"type\": \"object\""));
return errors;
if !has_combinators {
errors.push(format!("{path}: missing \"type\": \"object\""));
return errors;
}
}
}
// Rule 2: must have "properties" as an object
// Validate combinator variants recursively
for key in ["allOf", "oneOf", "anyOf"] {
if let Some(variants) = schema.get(key).and_then(|v| v.as_array()) {
for (i, variant) in variants.iter().enumerate() {
if variant.get("type").and_then(|t| t.as_str()) == Some("object")
|| variant.get("properties").is_some()
{
let variant_path = format!("{path}.{key}[{i}]");
errors.extend(check_object_schema(variant, &variant_path));
}
}
}
}
// Rule 2: must have "properties" as an object (unless combinators define them)
let properties = match schema.get("properties").and_then(|p| p.as_object()) {
Some(p) => p,
None => {
errors.push(format!("{path}: missing or non-object \"properties\""));
if !has_combinators {
errors.push(format!("{path}: missing or non-object \"properties\""));
return errors;
}
// Combinators define the structure — validate top-level `required` keys
// against merged properties from all combinator variants.
if let Some(required) = schema.get("required").and_then(|r| r.as_array()) {
let mut merged_keys = std::collections::HashSet::new();
if let Some(all_of) = schema.get("allOf").and_then(|a| a.as_array()) {
for variant in all_of {
if let Some(props) = variant.get("properties").and_then(|p| p.as_object()) {
merged_keys.extend(props.keys().cloned());
}
}
}
for key in ["oneOf", "anyOf"] {
if let Some(variants) = schema.get(key).and_then(|v| v.as_array()) {
for variant in variants {
if let Some(props) =
variant.get("properties").and_then(|p| p.as_object())
{
merged_keys.extend(props.keys().cloned());
}
}
}
}
for req in required {
if let Some(key) = req.as_str()
&& !merged_keys.contains(key)
{
errors.push(format!(
"{path}: required key \"{key}\" not found in any combinator variant properties"
));
}
}
}
return errors;
}
};
+127 -5
View File
@@ -1,5 +1,6 @@
//! Tool trait and types.
use std::fmt;
use std::time::Duration;
use async_trait::async_trait;
@@ -112,6 +113,33 @@ impl Default for ToolRateLimitConfig {
}
}
/// Risk level of a tool invocation.
///
/// Used by the shell tool to classify commands and by the worker to drive
/// approval decisions and observability logging. Implements `Ord` so callers
/// can compare levels (e.g. `risk >= RiskLevel::High`).
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
pub enum RiskLevel {
/// Read-only, safe, reversible (e.g. `ls`, `cat`, `grep`).
Low,
/// Creates or modifies state, but generally reversible
/// (e.g. `mkdir`, `git commit`, `cargo build`).
Medium,
/// Destructive, irreversible, or security-sensitive
/// (e.g. `rm -rf`, `git push --force`, `kill -9`).
High,
}
impl fmt::Display for RiskLevel {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Low => f.write_str("low"),
Self::Medium => f.write_str("medium"),
Self::High => f.write_str("high"),
}
}
}
/// Where a tool should execute: orchestrator process or inside a container.
///
/// Orchestrator tools run in the main agent process (memory access, job mgmt, etc).
@@ -276,6 +304,18 @@ pub trait Tool: Send + Sync {
true
}
/// Risk level for a specific invocation of this tool.
///
/// Defaults to `Low` (read-only, safe). Override for tools whose risk
/// depends on the parameters — the shell tool classifies commands into
/// `Low` / `Medium` / `High` based on the command string.
///
/// The worker logs this value with every tool call so operators can audit
/// the risk level at which each execution was classified.
fn risk_level_for(&self, _params: &serde_json::Value) -> RiskLevel {
RiskLevel::Low
}
/// Whether this tool invocation requires user approval.
///
/// Returns `Never` by default (most tools run in a sandboxed environment).
@@ -462,6 +502,22 @@ pub fn redact_params(params: &serde_json::Value, sensitive: &[&str]) -> serde_js
/// on maliciously crafted schemas.
const MAX_SCHEMA_DEPTH: usize = 16;
/// Returns true if the schema uses `oneOf`, `anyOf`, or `allOf` combinators
/// where at least one variant is an object type (has `type: "object"` or `properties`).
fn has_object_combinator_variants(schema: &serde_json::Value) -> bool {
for key in ["oneOf", "anyOf", "allOf"] {
if let Some(variants) = schema.get(key).and_then(|v| v.as_array())
&& variants.iter().any(|v| {
v.get("type").and_then(|t| t.as_str()) == Some("object")
|| v.get("properties").is_some()
})
{
return true;
}
}
false
}
pub fn validate_tool_schema(schema: &serde_json::Value, path: &str) -> Vec<String> {
validate_tool_schema_inner(schema, path, 0)
}
@@ -476,7 +532,18 @@ fn validate_tool_schema_inner(schema: &serde_json::Value, path: &str, depth: usi
return errors;
}
// Rule 1: must have "type": "object" at this level
// Report non-array combinator values as errors.
for key in ["oneOf", "anyOf", "allOf"] {
if let Some(val) = schema.get(key)
&& !val.is_array()
{
errors.push(format!("{path}: \"{key}\" must be an array"));
}
}
let has_combinators = has_object_combinator_variants(schema);
// Rule 1: must have "type": "object" at this level (unless combinators define the structure)
match schema.get("type").and_then(|t| t.as_str()) {
Some("object") => {}
Some(other) => {
@@ -484,16 +551,71 @@ fn validate_tool_schema_inner(schema: &serde_json::Value, path: &str, depth: usi
return errors; // Can't check further
}
None => {
errors.push(format!("{path}: missing \"type\": \"object\""));
return errors;
if !has_combinators {
errors.push(format!("{path}: missing \"type\": \"object\""));
return errors;
}
}
}
// Rule 2: must have "properties" as an object
// Validate combinator variants recursively
for key in ["allOf", "oneOf", "anyOf"] {
if let Some(variants) = schema.get(key).and_then(|v| v.as_array()) {
for (i, variant) in variants.iter().enumerate() {
if variant.get("type").and_then(|t| t.as_str()) == Some("object")
|| variant.get("properties").is_some()
{
let variant_path = format!("{path}.{key}[{i}]");
errors.extend(validate_tool_schema_inner(
variant,
&variant_path,
depth + 1,
));
}
}
}
}
// Rule 2: must have "properties" as an object (unless combinators define them)
let properties = match schema.get("properties").and_then(|p| p.as_object()) {
Some(p) => p,
None => {
errors.push(format!("{path}: missing or non-object \"properties\""));
if !has_combinators {
errors.push(format!("{path}: missing or non-object \"properties\""));
return errors;
}
// Combinators define the structure — validate top-level `required` keys
// against merged properties from all combinator variants.
if let Some(required) = schema.get("required").and_then(|r| r.as_array()) {
let mut merged_keys = std::collections::HashSet::new();
if let Some(all_of) = schema.get("allOf").and_then(|a| a.as_array()) {
for variant in all_of {
if let Some(props) = variant.get("properties").and_then(|p| p.as_object()) {
merged_keys.extend(props.keys().cloned());
}
}
}
for key in ["oneOf", "anyOf"] {
if let Some(variants) = schema.get(key).and_then(|v| v.as_array()) {
for variant in variants {
if let Some(props) =
variant.get("properties").and_then(|p| p.as_object())
{
merged_keys.extend(props.keys().cloned());
}
}
}
}
for req in required {
if let Some(key) = req.as_str()
&& !merged_keys.contains(key)
{
errors.push(format!(
"{path}: required key \"{key}\" not found in any combinator variant properties"
));
}
}
}
return errors;
}
};
+99
View File
@@ -708,6 +708,9 @@ pub struct ToolSetupSchema {
/// Secrets the user must provide before the tool can be used.
#[serde(default)]
pub required_secrets: Vec<ToolSecretSetupSchema>,
/// Non-secret fields the user can configure in the setup modal.
#[serde(default)]
pub required_fields: Vec<ToolFieldSetupSchema>,
}
/// A single secret required during tool setup.
@@ -722,6 +725,46 @@ pub struct ToolSecretSetupSchema {
pub optional: bool,
}
/// A non-secret field required during tool setup.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ToolFieldSetupSchema {
/// Field name in setup payload.
pub name: String,
/// User-facing prompt shown in the setup modal.
pub prompt: String,
/// If true, the user may skip this field.
#[serde(default)]
pub optional: bool,
/// Input type used in the setup modal.
#[serde(default = "default_tool_setup_field_input_type")]
pub input_type: ToolSetupFieldInputType,
/// Optional dotted setting path to persist this value to.
///
/// Restricted by the host to extension-owned namespaces and a small
/// allowlist of approved global settings.
///
/// Example: `extensions.switch-llm.provider`, `llm_backend`, or
/// `selected_model`.
#[serde(default)]
pub setting_path: Option<String>,
/// Whether changing this field requires a restart to fully apply.
#[serde(default)]
pub restart_required: bool,
}
/// Input widget type for a setup field.
#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum ToolSetupFieldInputType {
#[default]
Text,
Password,
}
fn default_tool_setup_field_input_type() -> ToolSetupFieldInputType {
ToolSetupFieldInputType::Text
}
#[cfg(test)]
mod tests {
use crate::tools::wasm::capabilities_schema::{CapabilitiesFile, CredentialLocationSchema};
@@ -1218,6 +1261,20 @@ mod tests {
"prompt": "Google OAuth Client Secret",
"optional": true
}
],
"required_fields": [
{
"name": "llm_backend",
"prompt": "LLM Provider",
"setting_path": "llm_backend",
"restart_required": true
},
{
"name": "selected_model",
"prompt": "Model Name",
"input_type": "text",
"setting_path": "selected_model"
}
]
}
}"#;
@@ -1230,6 +1287,48 @@ mod tests {
assert!(!setup.required_secrets[0].optional);
assert_eq!(setup.required_secrets[1].name, "google_oauth_client_secret");
assert!(setup.required_secrets[1].optional);
assert_eq!(setup.required_fields.len(), 2);
assert_eq!(setup.required_fields[0].name, "llm_backend");
assert_eq!(
setup.required_fields[0].setting_path.as_deref(),
Some("llm_backend")
);
assert!(setup.required_fields[0].restart_required);
assert_eq!(
setup.required_fields[0].input_type,
crate::tools::wasm::capabilities_schema::ToolSetupFieldInputType::Text
);
assert_eq!(setup.required_fields[1].name, "selected_model");
}
#[test]
fn test_tool_setup_field_input_type_defaults_to_text() {
let json = r#"{
"setup": {
"required_fields": [
{
"name": "provider",
"prompt": "Provider"
},
{
"name": "token_hint",
"prompt": "Token Hint",
"input_type": "password"
}
]
}
}"#;
let caps = CapabilitiesFile::from_json(json).unwrap();
let setup = caps.setup.unwrap();
assert_eq!(
setup.required_fields[0].input_type,
crate::tools::wasm::capabilities_schema::ToolSetupFieldInputType::Text
);
assert_eq!(
setup.required_fields[1].input_type,
crate::tools::wasm::capabilities_schema::ToolSetupFieldInputType::Password
);
}
#[test]
+2 -2
View File
@@ -206,7 +206,7 @@ impl WasmToolLoader {
})
.await?;
tracing::info!(
tracing::debug!(
name = name,
wasm_path = %wasm_path.display(),
"Loaded WASM tool from file"
@@ -306,7 +306,7 @@ impl WasmToolLoader {
}
if !results.loaded.is_empty() {
tracing::info!(
tracing::debug!(
count = results.loaded.len(),
tools = ?results.loaded,
"Loaded WASM tools from directory"
+1 -1
View File
@@ -139,5 +139,5 @@ pub use loader::{
// Capabilities schema (for parsing *.capabilities.json files)
pub use capabilities_schema::{
AuthCapabilitySchema, CapabilitiesFile, OAuthConfigSchema, RateLimitSchema,
ValidationEndpointSchema,
ToolFieldSetupSchema, ToolSetupFieldInputType, ToolSetupSchema, ValidationEndpointSchema,
};
+1 -1
View File
@@ -312,7 +312,7 @@ impl WasmToolRuntime {
.insert(prepared.name.clone(), Arc::clone(&prepared));
}
tracing::info!(
tracing::debug!(
name = %prepared.name,
"Prepared WASM tool for execution"
);
+184 -19
View File
@@ -17,6 +17,7 @@ use wasmtime::component::Linker;
use wasmtime_wasi::{ResourceTable, WasiCtx, WasiCtxBuilder, WasiView};
use crate::context::JobContext;
use crate::llm::recording::{HttpExchangeRequest, HttpExchangeResponse, HttpInterceptor};
use crate::safety::LeakDetector;
use crate::secrets::SecretsStore;
use crate::tools::tool::{Tool, ToolError, ToolOutput};
@@ -99,6 +100,9 @@ struct StoreData {
/// Dedicated tokio runtime for HTTP requests, lazily initialized.
/// Reused across multiple `http_request` calls within one execution.
http_runtime: Option<tokio::runtime::Runtime>,
/// Optional HTTP interceptor for testing — returns canned responses
/// instead of making real requests when set.
http_interceptor: Option<Arc<dyn HttpInterceptor>>,
}
impl StoreData {
@@ -119,6 +123,7 @@ impl StoreData {
credentials,
host_credentials,
http_runtime: None,
http_interceptor: None,
}
}
@@ -344,6 +349,59 @@ impl near::agent::host::Host for StoreData {
);
}
let rt = self.http_runtime.as_ref().expect("just initialized"); // safety: is_none branch above guarantees Some
// If an HTTP interceptor is set (testing), short-circuit with a canned response.
if let Some(interceptor) = &self.http_interceptor {
let interceptor = Arc::clone(interceptor);
let intercept_url = url.clone();
let intercept_method = method.clone();
let mut intercept_headers: Vec<(String, String)> = headers
.iter()
.map(|(k, v)| (k.clone(), v.clone()))
.collect();
intercept_headers.sort_by(|a, b| a.0.cmp(&b.0));
let intercept_body = body
.as_ref()
.map(|b| String::from_utf8_lossy(b).to_string());
let intercepted = rt.block_on(async {
let req = HttpExchangeRequest {
method: intercept_method,
url: intercept_url,
headers: intercept_headers,
body: intercept_body,
};
interceptor.before_request(&req).await
});
if let Some(resp) = intercepted {
let resp_headers: HashMap<String, String> = resp
.headers
.iter()
.map(|(k, v)| (k.clone(), v.clone()))
.collect();
let resp_headers_json =
serde_json::to_string(&resp_headers).unwrap_or_else(|_| "{}".to_string());
return Ok(near::agent::host::HttpResponse {
status: resp.status,
headers_json: resp_headers_json,
body: resp.body.into_bytes(),
});
}
}
// Capture request metadata before headers/body are consumed by the reqwest
// builder. Used for after_response callback when a recording interceptor is set.
let interceptor_req = self.http_interceptor.as_ref().map(|_| HttpExchangeRequest {
method: method.clone(),
url: url.clone(),
headers: headers
.iter()
.map(|(k, v)| (k.clone(), v.clone()))
.collect(),
body: body
.as_ref()
.map(|b| String::from_utf8_lossy(b).to_string()),
});
let result = rt.block_on(async {
let client = reqwest::Client::builder()
.connect_timeout(Duration::from_secs(10))
@@ -434,6 +492,51 @@ impl near::agent::host::Host for StoreData {
})
});
// Notify the interceptor about the completed response (recording mode).
// RecordingHttpInterceptor returns None from before_request and captures
// exchanges via after_response, so this path is exercised during trace recording.
if let (Some(interceptor), Some(req), Ok(resp)) =
(&self.http_interceptor, &interceptor_req, &result)
{
let interceptor = Arc::clone(interceptor);
// Redact credentials from request before passing to the interceptor
// to prevent credential leakage into recorded traces.
let mut redacted_req = req.clone();
redacted_req.url = self.redact_credentials(&redacted_req.url);
redacted_req.headers = redacted_req
.headers
.into_iter()
.map(|(k, v)| (k, self.redact_credentials(&v)))
.collect();
redacted_req.body = redacted_req.body.map(|b| self.redact_credentials(&b));
let resp_headers: Vec<(String, String)> =
serde_json::from_str::<HashMap<String, String>>(&resp.headers_json)
.unwrap_or_default()
.into_iter()
.collect();
let resp_body = String::from_utf8_lossy(&resp.body).to_string();
// Redact credentials from response as well
let redacted_headers: Vec<(String, String)> = resp_headers
.into_iter()
.map(|(k, v)| (k, self.redact_credentials(&v)))
.collect();
let redacted_body = self.redact_credentials(&resp_body);
let exchange_resp = HttpExchangeResponse {
status: resp.status,
headers: redacted_headers,
body: redacted_body,
};
rt.block_on(async {
interceptor
.after_response(&redacted_req, &exchange_resp)
.await;
});
}
// Redact credentials from error messages before returning to WASM
result.map_err(|e| self.redact_credentials(&e))
}
@@ -476,6 +579,9 @@ pub struct WasmToolWrapper {
secrets_store: Option<Arc<dyn SecretsStore + Send + Sync>>,
/// OAuth refresh configuration for auto-refreshing expired tokens.
oauth_refresh: Option<OAuthRefreshConfig>,
/// Optional HTTP interceptor for testing — returns canned responses
/// instead of making real requests when set.
http_interceptor: Option<Arc<dyn HttpInterceptor>>,
}
#[derive(Debug, Clone)]
@@ -502,23 +608,51 @@ impl WasmToolSchemas {
}
fn is_permissive_schema(schema: &serde_json::Value) -> bool {
schema
if schema
.get("properties")
.and_then(|p| p.as_object())
.is_none_or(|p| p.is_empty())
.is_some_and(|p| !p.is_empty())
{
return false;
}
// Schemas with combinator variants containing properties are not permissive
for key in ["oneOf", "anyOf", "allOf"] {
if let Some(variants) = schema.get(key).and_then(|v| v.as_array())
&& variants.iter().any(|v| {
v.get("properties")
.and_then(|p| p.as_object())
.is_some_and(|p| !p.is_empty())
})
{
return false;
}
}
true
}
fn typed_property_count(schema: &serde_json::Value) -> usize {
schema
.get("properties")
.and_then(|p| p.as_object())
.map(|props| {
props
.values()
.filter(|prop| schema_is_typed_property(prop))
.count()
})
.unwrap_or(0)
let mut all_props = serde_json::Map::new();
if let Some(props) = schema.get("properties").and_then(|p| p.as_object()) {
all_props.extend(props.iter().map(|(k, v)| (k.clone(), v.clone())));
}
for key in ["allOf", "oneOf", "anyOf"] {
if let Some(variants) = schema.get(key).and_then(|v| v.as_array()) {
for variant in variants {
if let Some(props) = variant.get("properties").and_then(|p| p.as_object()) {
all_props.extend(props.iter().map(|(k, v)| (k.clone(), v.clone())));
}
}
}
}
all_props
.values()
.filter(|prop| schema_is_typed_property(prop))
.count()
}
fn new(discovery: serde_json::Value) -> Self {
@@ -564,9 +698,20 @@ impl WasmToolWrapper {
credentials: HashMap::new(),
secrets_store: None,
oauth_refresh: None,
http_interceptor: None,
}
}
/// Set an HTTP interceptor for testing.
///
/// When set, WASM tool HTTP requests are routed through the interceptor
/// instead of making real network calls. This allows tests to verify the
/// exact HTTP requests a WASM tool constructs.
pub fn with_http_interceptor(mut self, interceptor: Arc<dyn HttpInterceptor>) -> Self {
self.http_interceptor = Some(interceptor);
self
}
/// Override the tool description.
pub fn with_description(mut self, description: impl Into<String>) -> Self {
self.description = description.into();
@@ -651,12 +796,13 @@ impl WasmToolWrapper {
let limits = &self.prepared.limits;
// Create store with fresh state (NEAR pattern: fresh instance per call)
let store_data = StoreData::new(
let mut store_data = StoreData::new(
limits.memory_bytes,
self.capabilities.clone(),
self.credentials.clone(),
host_credentials,
);
store_data.http_interceptor = self.http_interceptor.clone();
let mut store = Store::new(engine, store_data);
// Configure fuel if enabled
@@ -872,6 +1018,7 @@ impl Tool for WasmToolWrapper {
credentials,
secrets_store: None, // Not needed in blocking task
oauth_refresh: None, // Already used above for pre-refresh
http_interceptor: self.http_interceptor.clone(),
};
tokio::task::spawn_blocking(move || {
@@ -1320,15 +1467,33 @@ fn is_private_ip(ip: std::net::IpAddr) -> bool {
}
fn schema_contains_container_properties(schema: &serde_json::Value) -> bool {
schema
let has_container = |props: &serde_json::Map<String, serde_json::Value>| {
props
.values()
.any(|prop| schema_declares_type(prop, "array") || schema_declares_type(prop, "object"))
};
if schema
.get("properties")
.and_then(|p| p.as_object())
.map(|props| {
props.values().any(|prop| {
schema_declares_type(prop, "array") || schema_declares_type(prop, "object")
.is_some_and(has_container)
{
return true;
}
for key in ["allOf", "oneOf", "anyOf"] {
if let Some(variants) = schema.get(key).and_then(|v| v.as_array())
&& variants.iter().any(|v| {
v.get("properties")
.and_then(|p| p.as_object())
.is_some_and(has_container)
})
})
.unwrap_or(false)
{
return true;
}
}
false
}
fn schema_declares_type(schema: &serde_json::Value, expected: &str) -> bool {
+3 -3
View File
@@ -190,7 +190,7 @@ pub async fn start_managed_tunnel(
mut config: crate::config::Config,
) -> (crate::config::Config, Option<Box<dyn Tunnel>>) {
if config.tunnel.public_url.is_some() {
tracing::info!(
tracing::debug!(
"Static tunnel URL in use: {}",
config.tunnel.public_url.as_deref().unwrap_or("?")
);
@@ -216,7 +216,7 @@ pub async fn start_managed_tunnel(
match create_tunnel(provider_config) {
Ok(Some(tunnel)) => {
tracing::info!(
tracing::debug!(
"Starting {} tunnel on {}:{}...",
tunnel.name(),
gateway_host,
@@ -224,7 +224,7 @@ pub async fn start_managed_tunnel(
);
match tunnel.start(gateway_host, gateway_port).await {
Ok(url) => {
tracing::info!("Tunnel started: {}", url);
tracing::debug!("Tunnel started: {}", url);
config.tunnel.public_url = Some(url);
(config, Some(tunnel))
}
+1 -1
View File
@@ -472,7 +472,7 @@ impl LoopDelegate for ContainerDelegate {
"tool_name": tc.name,
"output": match &result {
Ok(output) => truncate_for_preview(output, 2000),
Err(e) => format!("Error: {}", truncate_for_preview(e, 500)),
Err(e) => format!("Error: {}", truncate_for_preview(e, 500)).into(),
},
"success": result.is_ok(),
}),
+7 -1
View File
@@ -592,10 +592,12 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
// Redact sensitive parameter values before they touch any observability or audit path.
let safe_params = redact_params(&effective_params, tool.sensitive_params());
let risk = tool.risk_level_for(&effective_params);
tracing::debug!(
tool = %tool_name,
params = %safe_params,
job = %job_id,
risk = %risk,
"Tool call started"
);
@@ -798,12 +800,16 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
});
}
let error_preview = {
let msg = format!("Error: {}", e);
truncate_for_preview(&msg, 500).into_owned()
};
self.log_event(
"tool_result",
serde_json::json!({
"tool_name": selection.tool_name,
"success": false,
"output": truncate_for_preview(&format!("Error: {}", e), 500),
"output": error_preview,
}),
);
+4
View File
@@ -558,6 +558,10 @@ impl Workspace {
/// which uses `\n\n`.
pub async fn append(&self, path: &str, content: &str) -> Result<(), WorkspaceError> {
let path = normalize_path(path);
// Scan system-prompt-injected files for prompt injection.
if is_system_prompt_file(&path) && !content.is_empty() {
reject_if_injected(&path, content)?;
}
let doc = self
.storage
.get_or_create_document_by_path(&self.user_id, self.agent_id, &path)