mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-09-03 01:59:23 +00:00
feat: 10 infrastructure improvements from zeroclaw (#126)
* refactor: break up agent_loop.rs into four focused modules Split the monolithic 2835-line agent_loop.rs into: - agent_loop.rs (722L): Agent struct, event loop, message dispatch - dispatcher.rs (635L): Agentic tool loop, tool execution, auth detection - commands.rs (484L): System commands, job handlers, heartbeat, summarize - thread_ops.rs (1059L): Thread lifecycle, approval, undo/redo, persistence Each module gets its own impl Agent block. Agent fields changed to pub(super) so sibling modules in the agent package can access them. All 16 existing tests pass in their new locations. Inspired by ZeroClaw's agent module split (agent.rs, loop_.rs, dispatcher.rs, prompt.rs, memory_loader.rs). Co-Authored-By: Claude Opus 4.6 <[email protected]> * feat: add cost caps and guardrails for autonomous agent spending Daily budget (MAX_COST_PER_DAY_CENTS) and hourly action rate (MAX_ACTIONS_PER_HOUR) limits prevent runaway agents from burning through API credits, especially in daemon/heartbeat modes. - CostGuard with pre-flight check and post-call recording - Sliding window for hourly rate, midnight-UTC daily reset - 80% threshold warning, atomic fast-path for exceeded budget - Wired into dispatcher loop (check before LLM call, record after) Co-Authored-By: Claude Opus 4.6 <[email protected]> * feat: add circuit breaker on LLM providers Wraps LlmProvider with a Closed/Open/HalfOpen state machine that trips after consecutive transient failures, preventing request storms against a degraded backend. Automatically probes for recovery. - CircuitBreakerProvider implements LlmProvider (drop-in wrapper) - Transient error classification (server, rate-limit, network, auth infra) - Client errors (wrong model, context overflow) don't trip the breaker - Configurable via CIRCUIT_BREAKER_THRESHOLD and CIRCUIT_BREAKER_RECOVERY_SECS - Composes with existing FailoverProvider (circuit breaker wraps failover) - 12 tests covering full state machine and error classification Co-Authored-By: Claude Opus 4.6 <[email protected]> * feat: add tunnel abstraction for remote access Trait-based tunnel system with lifecycle management (start/stop/health) for exposing the agent to the internet through external tunnel binaries. Five providers: - Cloudflare Tunnel (cloudflared, Zero Trust token auth) - Tailscale (serve for tailnet, funnel for public) - ngrok (with optional custom domain) - Custom (arbitrary command with {host}/{port} placeholders) - None (local-only, no external exposure) Config via TUNNEL_PROVIDER + provider-specific env vars. Extends existing TunnelConfig with optional managed provider alongside the static TUNNEL_URL path. Factory, shared process management, and 37 tests covering all providers and edge cases. Co-Authored-By: Claude Opus 4.6 <[email protected]> * feat: add OS service management (launchd/systemd) Adds `ironclaw service {install,start,stop,status,uninstall}` for running the agent as a background daemon. macOS uses launchd plists under ~/Library/LaunchAgents, Linux uses systemd user units. Co-Authored-By: Claude Opus 4.6 <[email protected]> * feat: add observability trait system with noop, log, and multi backends Introduces an Observer trait for recording agent lifecycle events and metrics, with pluggable backends. The noop backend compiles to zero overhead, log backend uses tracing, and multi fans out to multiple observers. Configured via OBSERVABILITY_BACKEND env var. Co-Authored-By: Claude Opus 4.6 <[email protected]> * feat: add in-memory LLM response cache with TTL and LRU eviction CachedProvider wraps any LlmProvider and caches complete() responses keyed by SHA-256(model + messages). Tool-calling requests are never cached since they trigger side effects. Configurable via RESPONSE_CACHE_ENABLED, RESPONSE_CACHE_TTL_SECS, and RESPONSE_CACHE_MAX_ENTRIES env vars. Co-Authored-By: Claude Opus 4.6 <[email protected]> * feat: add memory hygiene with cadence-gated daily log cleanup Adds workspace::hygiene module that automatically deletes daily log documents older than a configurable retention period (default 30 days). Runs on a 12-hour cadence tracked via a local state file to avoid redundant passes. Best-effort design: failures are logged, never fatal. Co-Authored-By: Claude Opus 4.6 <[email protected]> * feat: add doctor diagnostics command for active health probing Probes external dependencies (Docker, cloudflared, ngrok, tailscale), validates NEAR AI session, checks database connectivity, and verifies workspace directory. Complements the passive `status` command. Co-Authored-By: Claude Opus 4.6 <[email protected]> * feat: add structured TOML config file support Adds ~/.ironclaw/config.toml as a configuration layer between env vars and database settings. Priority: env var > TOML file > DB > defaults. - `ironclaw config init` generates a commented config.toml from current settings - `ironclaw --config path/to/config.toml` loads a custom config file - Settings.merge_from() only overlays non-default values from the TOML file - `ironclaw config path` now shows TOML file status Co-Authored-By: Claude Opus 4.6 <[email protected]> * fix: address codex review findings - apply_toml_overlay now returns Result and errors on explicit missing or invalid config paths (was log-only, violating the documented contract that explicit paths are fatal) - custom tunnel url_pattern is now used to filter extracted URLs, not just as a gate for scanning stdout - systemd ExecStart path is now quoted to handle spaces in paths Co-Authored-By: Claude Opus 4.6 <[email protected]> * fix: address PR review feedback - Cache key now includes max_tokens, temperature, and stop_sequences so different request parameters produce distinct keys - to_cents() uses .trunc() + parse::<u64> instead of f64 intermediary, avoiding precision loss for large values - Tailscale public URL no longer includes local port (serve/funnel expose on standard HTTPS port 443) Co-Authored-By: Claude Opus 4.6 <[email protected]> * feat: wire up tunnel lifecycle and fix audit findings Connect the tunnel module to the rest of the application so that setting TUNNEL_PROVIDER actually starts a managed tunnel at boot and stops it on shutdown. Previously create_tunnel() was never called outside tests. Changes: - Expand TunnelSettings with provider credential fields (settings.rs) - TunnelConfig::resolve() falls back to DB settings when env vars unset - Start tunnel at boot, stop on shutdown, show URL in boot screen - Setup wizard collects provider-specific credentials (ngrok, cloudflare, tailscale, custom, static URL) - Fix public_url() returning None under lock contention (SharedUrl) - Fix local_host parameter ignored by cloudflare/ngrok/tailscale - Fix tailscale silent fallback to "localhost" on bad JSON - Fix ngrok globally mutating config via add-authtoken (use env var) - Add 10s timeout to tailscale status --json Co-Authored-By: Claude Opus 4.6 <[email protected]> * fix: address PR review comments - Document split_whitespace limitation in CustomTunnel doc comment - Remove unnecessary quotes from systemd ExecStart directive Co-Authored-By: Claude Opus 4.6 <[email protected]> * fix: address PR review feedback (round 3) - doctor: missing libSQL DB on fresh install is Pass, not Fail - service: quote ExecStart path for systemd space handling Co-Authored-By: Claude Opus 4.6 <[email protected]> * fix: correct cost guard doc comment (LLM calls, not LLM/tool) Co-Authored-By: Claude Opus 4.6 <[email protected]> --------- Co-authored-by: Claude Opus 4.6 <[email protected]>
This commit is contained in:
co-authored by
Claude Opus 4.6
parent
436dda0f2f
commit
a158eee1b0
@@ -0,0 +1,635 @@
|
||||
//! Tool dispatch logic for the agent.
|
||||
//!
|
||||
//! Extracted from `agent_loop.rs` to keep the core agentic tool execution
|
||||
//! loop (LLM call -> tool calls -> repeat) in its own focused module.
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use tokio::sync::Mutex;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::agent::Agent;
|
||||
use crate::agent::session::{PendingApproval, Session, ThreadState};
|
||||
use crate::channels::{IncomingMessage, StatusUpdate};
|
||||
use crate::context::JobContext;
|
||||
use crate::error::Error;
|
||||
use crate::llm::{ChatMessage, Reasoning, ReasoningContext, RespondResult};
|
||||
|
||||
/// Result of the agentic loop execution.
|
||||
pub(super) enum AgenticLoopResult {
|
||||
/// Completed with a response.
|
||||
Response(String),
|
||||
/// A tool requires approval before continuing.
|
||||
NeedApproval {
|
||||
/// The pending approval request to store.
|
||||
pending: PendingApproval,
|
||||
},
|
||||
}
|
||||
|
||||
impl Agent {
|
||||
/// Run the agentic loop: call LLM, execute tools, repeat until text response.
|
||||
///
|
||||
/// Returns `AgenticLoopResult::Response` on completion, or
|
||||
/// `AgenticLoopResult::NeedApproval` if a tool requires user approval.
|
||||
///
|
||||
/// When `resume_after_tool` is true the loop already knows a tool was
|
||||
/// executed earlier in this turn (e.g. an approved tool), so it won't
|
||||
/// force the LLM to use tools if it responds with text.
|
||||
pub(super) async fn run_agentic_loop(
|
||||
&self,
|
||||
message: &IncomingMessage,
|
||||
session: Arc<Mutex<Session>>,
|
||||
thread_id: Uuid,
|
||||
initial_messages: Vec<ChatMessage>,
|
||||
resume_after_tool: bool,
|
||||
) -> Result<AgenticLoopResult, Error> {
|
||||
// Load workspace system prompt (identity files: AGENTS.md, SOUL.md, etc.)
|
||||
let system_prompt = if let Some(ws) = self.workspace() {
|
||||
match ws.system_prompt().await {
|
||||
Ok(prompt) if !prompt.is_empty() => Some(prompt),
|
||||
Ok(_) => None,
|
||||
Err(e) => {
|
||||
tracing::debug!("Could not load workspace system prompt: {}", e);
|
||||
None
|
||||
}
|
||||
}
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let mut reasoning = Reasoning::new(self.llm().clone(), self.safety().clone());
|
||||
if let Some(prompt) = system_prompt {
|
||||
reasoning = reasoning.with_system_prompt(prompt);
|
||||
}
|
||||
|
||||
// Build context with messages that we'll mutate during the loop
|
||||
let mut context_messages = initial_messages;
|
||||
|
||||
// Create a JobContext for tool execution (chat doesn't have a real job)
|
||||
let job_ctx = JobContext::with_user(&message.user_id, "chat", "Interactive chat session");
|
||||
|
||||
const MAX_TOOL_ITERATIONS: usize = 10;
|
||||
let mut iteration = 0;
|
||||
let mut tools_executed = resume_after_tool;
|
||||
|
||||
loop {
|
||||
iteration += 1;
|
||||
if iteration > MAX_TOOL_ITERATIONS {
|
||||
return Err(crate::error::LlmError::InvalidResponse {
|
||||
provider: "agent".to_string(),
|
||||
reason: format!("Exceeded maximum tool iterations ({})", MAX_TOOL_ITERATIONS),
|
||||
}
|
||||
.into());
|
||||
}
|
||||
|
||||
// Check if interrupted
|
||||
{
|
||||
let sess = session.lock().await;
|
||||
if let Some(thread) = sess.threads.get(&thread_id)
|
||||
&& thread.state == ThreadState::Interrupted
|
||||
{
|
||||
return Err(crate::error::JobError::ContextError {
|
||||
id: thread_id,
|
||||
reason: "Interrupted".to_string(),
|
||||
}
|
||||
.into());
|
||||
}
|
||||
}
|
||||
|
||||
// Enforce cost guardrails before the LLM call
|
||||
if let Err(limit) = self.cost_guard().check_allowed().await {
|
||||
return Err(crate::error::LlmError::InvalidResponse {
|
||||
provider: "agent".to_string(),
|
||||
reason: limit.to_string(),
|
||||
}
|
||||
.into());
|
||||
}
|
||||
|
||||
// Refresh tool definitions each iteration so newly built tools become visible
|
||||
let tool_defs = self.tools().tool_definitions().await;
|
||||
|
||||
// Call LLM with current context
|
||||
let context = ReasoningContext::new()
|
||||
.with_messages(context_messages.clone())
|
||||
.with_tools(tool_defs)
|
||||
.with_metadata({
|
||||
let mut m = std::collections::HashMap::new();
|
||||
m.insert("thread_id".to_string(), thread_id.to_string());
|
||||
m
|
||||
});
|
||||
|
||||
let output = reasoning.respond_with_tools(&context).await?;
|
||||
|
||||
// Record cost and track token usage
|
||||
let model_name = self.llm().active_model_name();
|
||||
let call_cost = self
|
||||
.cost_guard()
|
||||
.record_llm_call(
|
||||
&model_name,
|
||||
output.usage.input_tokens,
|
||||
output.usage.output_tokens,
|
||||
)
|
||||
.await;
|
||||
tracing::debug!(
|
||||
"LLM call used {} input + {} output tokens (${:.6})",
|
||||
output.usage.input_tokens,
|
||||
output.usage.output_tokens,
|
||||
call_cost,
|
||||
);
|
||||
|
||||
match output.result {
|
||||
RespondResult::Text(text) => {
|
||||
// If no tools have been executed yet, prompt the LLM to use tools
|
||||
// This handles the case where the model explains what it will do
|
||||
// instead of actually calling tools
|
||||
if !tools_executed && iteration < 3 {
|
||||
tracing::debug!(
|
||||
"No tools executed yet (iteration {}), prompting for tool use",
|
||||
iteration
|
||||
);
|
||||
context_messages.push(ChatMessage::assistant(&text));
|
||||
context_messages.push(ChatMessage::user(
|
||||
"Please proceed and use the available tools to complete this task.",
|
||||
));
|
||||
continue;
|
||||
}
|
||||
|
||||
// Tools have been executed or we've tried multiple times, return response
|
||||
return Ok(AgenticLoopResult::Response(text));
|
||||
}
|
||||
RespondResult::ToolCalls {
|
||||
tool_calls,
|
||||
content,
|
||||
} => {
|
||||
tools_executed = true;
|
||||
|
||||
// Add the assistant message with tool_calls to context.
|
||||
// OpenAI protocol requires this before tool-result messages.
|
||||
context_messages.push(ChatMessage::assistant_with_tool_calls(
|
||||
content,
|
||||
tool_calls.clone(),
|
||||
));
|
||||
|
||||
// Execute tools and add results to context
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::Thinking(format!(
|
||||
"Executing {} tool(s)...",
|
||||
tool_calls.len()
|
||||
)),
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
|
||||
// Record tool calls in the thread
|
||||
{
|
||||
let mut sess = session.lock().await;
|
||||
if let Some(thread) = sess.threads.get_mut(&thread_id)
|
||||
&& let Some(turn) = thread.last_turn_mut()
|
||||
{
|
||||
for tc in &tool_calls {
|
||||
turn.record_tool_call(&tc.name, tc.arguments.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Execute each tool (with approval checking and hook interception)
|
||||
for mut tc in tool_calls {
|
||||
// Check if tool requires approval
|
||||
if let Some(tool) = self.tools().get(&tc.name).await
|
||||
&& tool.requires_approval()
|
||||
{
|
||||
// Check if auto-approved for this session
|
||||
let mut is_auto_approved = {
|
||||
let sess = session.lock().await;
|
||||
sess.is_tool_auto_approved(&tc.name)
|
||||
};
|
||||
|
||||
// Override auto-approval for destructive parameters
|
||||
// (e.g. `rm -rf`, `git push --force` in shell commands).
|
||||
if is_auto_approved && tool.requires_approval_for(&tc.arguments) {
|
||||
tracing::info!(
|
||||
tool = %tc.name,
|
||||
"Parameters require explicit approval despite auto-approve"
|
||||
);
|
||||
is_auto_approved = false;
|
||||
}
|
||||
|
||||
if !is_auto_approved {
|
||||
// Need approval - store pending request and return
|
||||
let pending = PendingApproval {
|
||||
request_id: Uuid::new_v4(),
|
||||
tool_name: tc.name.clone(),
|
||||
parameters: tc.arguments.clone(),
|
||||
description: tool.description().to_string(),
|
||||
tool_call_id: tc.id.clone(),
|
||||
context_messages: context_messages.clone(),
|
||||
};
|
||||
|
||||
return Ok(AgenticLoopResult::NeedApproval { pending });
|
||||
}
|
||||
}
|
||||
|
||||
// Hook: BeforeToolCall — allow hooks to modify or reject tool calls
|
||||
{
|
||||
let event = crate::hooks::HookEvent::ToolCall {
|
||||
tool_name: tc.name.clone(),
|
||||
parameters: tc.arguments.clone(),
|
||||
user_id: message.user_id.clone(),
|
||||
context: "chat".to_string(),
|
||||
};
|
||||
match self.hooks().run(&event).await {
|
||||
Err(crate::hooks::HookError::Rejected { reason }) => {
|
||||
context_messages.push(ChatMessage::tool_result(
|
||||
&tc.id,
|
||||
&tc.name,
|
||||
format!("Tool call rejected by hook: {}", reason),
|
||||
));
|
||||
continue;
|
||||
}
|
||||
Err(err) => {
|
||||
context_messages.push(ChatMessage::tool_result(
|
||||
&tc.id,
|
||||
&tc.name,
|
||||
format!("Tool call blocked by hook policy: {}", err),
|
||||
));
|
||||
continue;
|
||||
}
|
||||
Ok(crate::hooks::HookOutcome::Continue {
|
||||
modified: Some(new_params),
|
||||
}) => match serde_json::from_str(&new_params) {
|
||||
Ok(parsed) => tc.arguments = parsed,
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
tool = %tc.name,
|
||||
"Hook returned non-JSON modification for ToolCall, ignoring: {}",
|
||||
e
|
||||
);
|
||||
}
|
||||
},
|
||||
_ => {} // Continue, fail-open errors already logged
|
||||
}
|
||||
}
|
||||
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::ToolStarted {
|
||||
name: tc.name.clone(),
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
|
||||
let tool_result = self
|
||||
.execute_chat_tool(&tc.name, &tc.arguments, &job_ctx)
|
||||
.await;
|
||||
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::ToolCompleted {
|
||||
name: tc.name.clone(),
|
||||
success: tool_result.is_ok(),
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
|
||||
if let Ok(ref output) = tool_result
|
||||
&& !output.is_empty()
|
||||
{
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::ToolResult {
|
||||
name: tc.name.clone(),
|
||||
preview: output.clone(),
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
// Record result in thread
|
||||
{
|
||||
let mut sess = session.lock().await;
|
||||
if let Some(thread) = sess.threads.get_mut(&thread_id)
|
||||
&& let Some(turn) = thread.last_turn_mut()
|
||||
{
|
||||
match &tool_result {
|
||||
Ok(output) => {
|
||||
turn.record_tool_result(serde_json::json!(output));
|
||||
}
|
||||
Err(e) => {
|
||||
turn.record_tool_error(e.to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// If tool_auth returned awaiting_token, enter auth mode
|
||||
// and short-circuit: return the instructions directly so
|
||||
// the LLM doesn't get a chance to hallucinate tool calls.
|
||||
if let Some((ext_name, instructions)) =
|
||||
detect_auth_awaiting(&tc.name, &tool_result)
|
||||
{
|
||||
let auth_data = parse_auth_result(&tool_result);
|
||||
{
|
||||
let mut sess = session.lock().await;
|
||||
if let Some(thread) = sess.threads.get_mut(&thread_id) {
|
||||
thread.enter_auth_mode(ext_name.clone());
|
||||
}
|
||||
}
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::AuthRequired {
|
||||
extension_name: ext_name,
|
||||
instructions: Some(instructions.clone()),
|
||||
auth_url: auth_data.auth_url,
|
||||
setup_url: auth_data.setup_url,
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
return Ok(AgenticLoopResult::Response(instructions));
|
||||
}
|
||||
|
||||
// Add tool result to context for next LLM call
|
||||
let result_content = match tool_result {
|
||||
Ok(output) => {
|
||||
// Sanitize output before showing to LLM
|
||||
let sanitized =
|
||||
self.safety().sanitize_tool_output(&tc.name, &output);
|
||||
self.safety().wrap_for_llm(
|
||||
&tc.name,
|
||||
&sanitized.content,
|
||||
sanitized.was_modified,
|
||||
)
|
||||
}
|
||||
Err(e) => format!("Error: {}", e),
|
||||
};
|
||||
|
||||
context_messages.push(ChatMessage::tool_result(
|
||||
&tc.id,
|
||||
&tc.name,
|
||||
result_content,
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Execute a tool for chat (without full job context).
|
||||
pub(super) async fn execute_chat_tool(
|
||||
&self,
|
||||
tool_name: &str,
|
||||
params: &serde_json::Value,
|
||||
job_ctx: &JobContext,
|
||||
) -> Result<String, Error> {
|
||||
let tool =
|
||||
self.tools()
|
||||
.get(tool_name)
|
||||
.await
|
||||
.ok_or_else(|| crate::error::ToolError::NotFound {
|
||||
name: tool_name.to_string(),
|
||||
})?;
|
||||
|
||||
// Validate tool parameters
|
||||
let validation = self.safety().validator().validate_tool_params(params);
|
||||
if !validation.is_valid {
|
||||
let details = validation
|
||||
.errors
|
||||
.iter()
|
||||
.map(|e| format!("{}: {}", e.field, e.message))
|
||||
.collect::<Vec<_>>()
|
||||
.join("; ");
|
||||
return Err(crate::error::ToolError::InvalidParameters {
|
||||
name: tool_name.to_string(),
|
||||
reason: format!("Invalid tool parameters: {}", details),
|
||||
}
|
||||
.into());
|
||||
}
|
||||
|
||||
tracing::debug!(
|
||||
tool = %tool_name,
|
||||
params = %params,
|
||||
"Tool call started"
|
||||
);
|
||||
|
||||
// Execute with per-tool timeout
|
||||
let timeout = tool.execution_timeout();
|
||||
let start = std::time::Instant::now();
|
||||
let result = tokio::time::timeout(timeout, async {
|
||||
tool.execute(params.clone(), job_ctx).await
|
||||
})
|
||||
.await;
|
||||
let elapsed = start.elapsed();
|
||||
|
||||
match &result {
|
||||
Ok(Ok(output)) => {
|
||||
let result_str = serde_json::to_string(&output.result)
|
||||
.unwrap_or_else(|_| "<serialize error>".to_string());
|
||||
tracing::debug!(
|
||||
tool = %tool_name,
|
||||
elapsed_ms = elapsed.as_millis() as u64,
|
||||
result = %result_str,
|
||||
"Tool call succeeded"
|
||||
);
|
||||
}
|
||||
Ok(Err(e)) => {
|
||||
tracing::debug!(
|
||||
tool = %tool_name,
|
||||
elapsed_ms = elapsed.as_millis() as u64,
|
||||
error = %e,
|
||||
"Tool call failed"
|
||||
);
|
||||
}
|
||||
Err(_) => {
|
||||
tracing::debug!(
|
||||
tool = %tool_name,
|
||||
elapsed_ms = elapsed.as_millis() as u64,
|
||||
timeout_secs = timeout.as_secs(),
|
||||
"Tool call timed out"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
let result = result
|
||||
.map_err(|_| crate::error::ToolError::Timeout {
|
||||
name: tool_name.to_string(),
|
||||
timeout,
|
||||
})?
|
||||
.map_err(|e| crate::error::ToolError::ExecutionFailed {
|
||||
name: tool_name.to_string(),
|
||||
reason: e.to_string(),
|
||||
})?;
|
||||
|
||||
// Convert result to string
|
||||
serde_json::to_string_pretty(&result.result).map_err(|e| {
|
||||
crate::error::ToolError::ExecutionFailed {
|
||||
name: tool_name.to_string(),
|
||||
reason: format!("Failed to serialize result: {}", e),
|
||||
}
|
||||
.into()
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// Parsed auth result fields for emitting StatusUpdate::AuthRequired.
|
||||
pub(super) struct ParsedAuthData {
|
||||
pub(super) auth_url: Option<String>,
|
||||
pub(super) setup_url: Option<String>,
|
||||
}
|
||||
|
||||
/// Extract auth_url and setup_url from a tool_auth result JSON string.
|
||||
pub(super) fn parse_auth_result(result: &Result<String, Error>) -> ParsedAuthData {
|
||||
let parsed = result
|
||||
.as_ref()
|
||||
.ok()
|
||||
.and_then(|s| serde_json::from_str::<serde_json::Value>(s).ok());
|
||||
ParsedAuthData {
|
||||
auth_url: parsed
|
||||
.as_ref()
|
||||
.and_then(|v| v.get("auth_url"))
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string()),
|
||||
setup_url: parsed
|
||||
.as_ref()
|
||||
.and_then(|v| v.get("setup_url"))
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
/// Check if a tool_auth result indicates the extension is awaiting a token.
|
||||
///
|
||||
/// Returns `Some((extension_name, instructions))` if the tool result contains
|
||||
/// `awaiting_token: true`, meaning the thread should enter auth mode.
|
||||
pub(super) fn detect_auth_awaiting(
|
||||
tool_name: &str,
|
||||
result: &Result<String, Error>,
|
||||
) -> Option<(String, String)> {
|
||||
if tool_name != "tool_auth" && tool_name != "tool_activate" {
|
||||
return None;
|
||||
}
|
||||
let output = result.as_ref().ok()?;
|
||||
let parsed: serde_json::Value = serde_json::from_str(output).ok()?;
|
||||
if parsed.get("awaiting_token") != Some(&serde_json::Value::Bool(true)) {
|
||||
return None;
|
||||
}
|
||||
let name = parsed.get("name")?.as_str()?.to_string();
|
||||
let instructions = parsed
|
||||
.get("instructions")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("Please provide your API token/key.")
|
||||
.to_string();
|
||||
Some((name, instructions))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::error::Error;
|
||||
|
||||
use super::detect_auth_awaiting;
|
||||
|
||||
#[test]
|
||||
fn test_detect_auth_awaiting_positive() {
|
||||
let result: Result<String, Error> = Ok(serde_json::json!({
|
||||
"name": "telegram",
|
||||
"kind": "WasmTool",
|
||||
"awaiting_token": true,
|
||||
"status": "awaiting_token",
|
||||
"instructions": "Please provide your Telegram Bot API token."
|
||||
})
|
||||
.to_string());
|
||||
|
||||
let detected = detect_auth_awaiting("tool_auth", &result);
|
||||
assert!(detected.is_some());
|
||||
let (name, instructions) = detected.unwrap();
|
||||
assert_eq!(name, "telegram");
|
||||
assert!(instructions.contains("Telegram Bot API"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_detect_auth_awaiting_not_awaiting() {
|
||||
let result: Result<String, Error> = Ok(serde_json::json!({
|
||||
"name": "telegram",
|
||||
"kind": "WasmTool",
|
||||
"awaiting_token": false,
|
||||
"status": "authenticated"
|
||||
})
|
||||
.to_string());
|
||||
|
||||
assert!(detect_auth_awaiting("tool_auth", &result).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_detect_auth_awaiting_wrong_tool() {
|
||||
let result: Result<String, Error> = Ok(serde_json::json!({
|
||||
"name": "telegram",
|
||||
"awaiting_token": true,
|
||||
})
|
||||
.to_string());
|
||||
|
||||
assert!(detect_auth_awaiting("tool_list", &result).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_detect_auth_awaiting_error_result() {
|
||||
let result: Result<String, Error> =
|
||||
Err(crate::error::ToolError::NotFound { name: "x".into() }.into());
|
||||
assert!(detect_auth_awaiting("tool_auth", &result).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_detect_auth_awaiting_default_instructions() {
|
||||
let result: Result<String, Error> = Ok(serde_json::json!({
|
||||
"name": "custom_tool",
|
||||
"awaiting_token": true,
|
||||
"status": "awaiting_token"
|
||||
})
|
||||
.to_string());
|
||||
|
||||
let (_, instructions) = detect_auth_awaiting("tool_auth", &result).unwrap();
|
||||
assert_eq!(instructions, "Please provide your API token/key.");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_detect_auth_awaiting_tool_activate() {
|
||||
let result: Result<String, Error> = Ok(serde_json::json!({
|
||||
"name": "slack",
|
||||
"kind": "McpServer",
|
||||
"awaiting_token": true,
|
||||
"status": "awaiting_token",
|
||||
"instructions": "Provide your Slack Bot token."
|
||||
})
|
||||
.to_string());
|
||||
|
||||
let detected = detect_auth_awaiting("tool_activate", &result);
|
||||
assert!(detected.is_some());
|
||||
let (name, instructions) = detected.unwrap();
|
||||
assert_eq!(name, "slack");
|
||||
assert!(instructions.contains("Slack Bot"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_detect_auth_awaiting_tool_activate_not_awaiting() {
|
||||
let result: Result<String, Error> = Ok(serde_json::json!({
|
||||
"name": "slack",
|
||||
"tools_loaded": ["slack_post_message"],
|
||||
"message": "Activated"
|
||||
})
|
||||
.to_string());
|
||||
|
||||
assert!(detect_auth_awaiting("tool_activate", &result).is_none());
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user