From 3f135bdde9ccfcfec353cddc1140e167d177003b Mon Sep 17 00:00:00 2001 From: Illia Polosukhin Date: Thu, 19 Feb 2026 18:28:15 -0800 Subject: [PATCH] fix: persist turns after approval and add agent-level tests (#250) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix: persist turns after approval and add agent-level tests Port relevant changes from PR #112 that were not carried over to #237: - Add persist_turn calls in process_approval for the response, error, and auth-required paths. Previously, turns completed after tool approval were never persisted to DB — if the process crashed after approval the entire turn (user message + assistant response) was lost. - Add agent-level unit tests: StaticLlmProvider mock, make_test_agent helper, tests for auto-approval logic, destructive shell command detection, and PendingApproval backward-compatible deserialization (without deferred_tool_calls field). - Remove unused _thread_state binding in process_approval. Co-Authored-By: Claude Opus 4.6 * fix: address 14 audit findings in src/agent/ Audit of the agent module found 2 High, 7 Medium, 3 Low, and 2 Nit severity issues. This commit fixes all of them: High: - Remove 4 `.expect()` calls in session.rs (entry API, match, direct indexing, if-let) to eliminate panic paths in production - Add typed RoutineError enum replacing Result<_, String> across routine.rs, routine_engine.rs, and callers in history/store.rs and db/libsql/mod.rs Medium: - Sanitize routine names in path construction to prevent directory traversal (routine_engine.rs) - Log warnings for 5 silently-swallowed errors in scheduler.rs, compaction.rs, and worker.rs - Extract shared handle_auth_intercept helper to deduplicate auth interception in thread_ops.rs - Add session count warning threshold in session_manager.rs - Make FullJob stub degradation visible via warn-level log and prepended warning in output Low: - Restrict dead code visibility with #[cfg(test)] on 19 unused items in submission.rs, task.rs, and undo.rs - Narrow pub to pub(crate) on self_repair.rs builder methods - Remove TaskStatus from mod.rs re-exports (test-only type) Co-Authored-By: Claude Opus 4.6 * fix: address PR review comments - Reorder persist_turn before persist_response_chain so the conversation row exists before the metadata UPDATE runs - Add persist_response_chain call to handle_auth_intercept so auth-required paths preserve the response chain - Harden sanitize_routine_name to use allowlist (alphanumeric, dash, underscore) instead of denylist replacements - Fix stale active_thread ID in get_or_create_thread: fall back to create_thread() when the stored ID is missing from the map - Persist turn on approval rejection so user messages survive crashes after a tool is rejected Co-Authored-By: Claude Opus 4.6 --------- Co-authored-by: Claude Opus 4.6 --- src/agent/compaction.rs | 22 +++- src/agent/dispatcher.rs | 226 +++++++++++++++++++++++++++++++++-- src/agent/mod.rs | 2 +- src/agent/routine.rs | 51 ++++++-- src/agent/routine_engine.rs | 83 +++++++++---- src/agent/scheduler.rs | 14 ++- src/agent/self_repair.rs | 14 ++- src/agent/session.rs | 26 ++-- src/agent/session_manager.rs | 11 ++ src/agent/submission.rs | 8 ++ src/agent/task.rs | 7 ++ src/agent/thread_ops.rs | 141 +++++++++++++--------- src/agent/undo.rs | 4 + src/agent/worker.rs | 71 +++++++---- src/db/libsql/mod.rs | 8 +- src/error.rs | 43 +++++++ src/history/store.rs | 8 +- 17 files changed, 583 insertions(+), 156 deletions(-) diff --git a/src/agent/compaction.rs b/src/agent/compaction.rs index 22b0ea6a..6e9479b6 100644 --- a/src/agent/compaction.rs +++ b/src/agent/compaction.rs @@ -105,7 +105,16 @@ impl ContextCompactor { // Write to workspace if available let summary_written = if let Some(ws) = workspace { - self.write_summary_to_workspace(ws, &summary).await.is_ok() + match self.write_summary_to_workspace(ws, &summary).await { + Ok(()) => true, + Err(e) => { + tracing::warn!( + "Compaction summary write failed (turns will still be truncated): {}", + e + ); + false + } + } } else { false }; @@ -157,7 +166,16 @@ impl ContextCompactor { let content = format_turns_for_storage(old_turns); // Write to workspace - let written = self.write_context_to_workspace(ws, &content).await.is_ok(); + let written = match self.write_context_to_workspace(ws, &content).await { + Ok(()) => true, + Err(e) => { + tracing::warn!( + "Compaction context write failed (turns will still be truncated): {}", + e + ); + false + } + }; // Truncate thread.truncate_turns(keep_recent); diff --git a/src/agent/dispatcher.rs b/src/agent/dispatcher.rs index 2bd1c871..d2a60717 100644 --- a/src/agent/dispatcher.rs +++ b/src/agent/dispatcher.rs @@ -402,7 +402,7 @@ impl Agent { // 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) + check_auth_required(&tc.name, &tool_result) { let auth_data = parse_auth_result(&tool_result); { @@ -581,7 +581,7 @@ pub(super) fn parse_auth_result(result: &Result) -> ParsedAuthDat /// /// 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( +pub(super) fn check_auth_required( tool_name: &str, result: &Result, ) -> Option<(String, String)> { @@ -604,9 +604,213 @@ pub(super) fn detect_auth_awaiting( #[cfg(test)] mod tests { - use crate::error::Error; + use std::sync::Arc; + use std::time::Duration; - use super::detect_auth_awaiting; + use async_trait::async_trait; + use rust_decimal::Decimal; + + use crate::agent::agent_loop::{Agent, AgentDeps}; + use crate::agent::cost_guard::{CostGuard, CostGuardConfig}; + use crate::agent::session::Session; + use crate::channels::ChannelManager; + use crate::config::{AgentConfig, SafetyConfig, SkillsConfig}; + use crate::context::ContextManager; + use crate::error::Error; + use crate::hooks::HookRegistry; + use crate::llm::{ + CompletionRequest, CompletionResponse, FinishReason, LlmProvider, ToolCall, + ToolCompletionRequest, ToolCompletionResponse, + }; + use crate::safety::SafetyLayer; + use crate::tools::ToolRegistry; + + use super::check_auth_required; + + /// Minimal LLM provider for unit tests that always returns a static response. + struct StaticLlmProvider; + + #[async_trait] + impl LlmProvider for StaticLlmProvider { + fn model_name(&self) -> &str { + "static-mock" + } + + fn cost_per_token(&self) -> (Decimal, Decimal) { + (Decimal::ZERO, Decimal::ZERO) + } + + async fn complete( + &self, + _request: CompletionRequest, + ) -> Result { + Ok(CompletionResponse { + content: "ok".to_string(), + input_tokens: 0, + output_tokens: 0, + finish_reason: FinishReason::Stop, + response_id: None, + }) + } + + async fn complete_with_tools( + &self, + _request: ToolCompletionRequest, + ) -> Result { + Ok(ToolCompletionResponse { + content: Some("ok".to_string()), + tool_calls: Vec::new(), + input_tokens: 0, + output_tokens: 0, + finish_reason: FinishReason::Stop, + response_id: None, + }) + } + } + + /// Build a minimal `Agent` for unit testing (no DB, no workspace, no extensions). + fn make_test_agent() -> Agent { + let deps = AgentDeps { + store: None, + llm: Arc::new(StaticLlmProvider), + cheap_llm: None, + safety: Arc::new(SafetyLayer::new(&SafetyConfig { + max_output_length: 100_000, + injection_check_enabled: true, + })), + tools: Arc::new(ToolRegistry::new()), + workspace: None, + extension_manager: None, + skill_registry: None, + skills_config: SkillsConfig::default(), + hooks: Arc::new(HookRegistry::new()), + cost_guard: Arc::new(CostGuard::new(CostGuardConfig::default())), + }; + + Agent::new( + AgentConfig { + name: "test-agent".to_string(), + max_parallel_jobs: 1, + job_timeout: Duration::from_secs(60), + stuck_threshold: Duration::from_secs(60), + repair_check_interval: Duration::from_secs(30), + max_repair_attempts: 1, + use_planning: false, + session_idle_timeout: Duration::from_secs(300), + allow_local_tools: false, + max_cost_per_day_cents: None, + max_actions_per_hour: None, + }, + deps, + ChannelManager::new(), + None, + None, + None, + Some(Arc::new(ContextManager::new(1))), + None, + ) + } + + #[test] + fn test_make_test_agent_succeeds() { + // Verify that a test agent can be constructed without panicking. + let _agent = make_test_agent(); + } + + #[test] + fn test_auto_approved_tool_is_respected() { + let _agent = make_test_agent(); + let mut session = Session::new("user-1"); + session.auto_approve_tool("http"); + + // A non-shell tool that is auto-approved should be approved. + assert!(session.is_tool_auto_approved("http")); + // A tool that hasn't been auto-approved should not be. + assert!(!session.is_tool_auto_approved("shell")); + } + + #[test] + fn test_shell_destructive_command_requires_approval_for() { + // ShellTool::requires_approval_for should detect destructive commands. + // This exercises the same code path used inline in run_agentic_loop. + use crate::tools::builtin::shell::requires_explicit_approval; + + let destructive_cmds = [ + "rm -rf /tmp/test", + "git push --force origin main", + "git reset --hard HEAD~5", + ]; + for cmd in &destructive_cmds { + assert!( + requires_explicit_approval(cmd), + "'{}' should require explicit approval", + cmd + ); + } + + 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 + ); + } + } + + #[test] + fn test_pending_approval_serialization_backcompat_without_deferred_calls() { + // PendingApproval from before the deferred_tool_calls field was added + // should deserialize with an empty vec (via #[serde(default)]). + let json = serde_json::json!({ + "request_id": uuid::Uuid::new_v4(), + "tool_name": "http", + "parameters": {"url": "https://example.com", "method": "GET"}, + "description": "Make HTTP request", + "tool_call_id": "call_123", + "context_messages": [{"role": "user", "content": "go"}] + }) + .to_string(); + + let parsed: crate::agent::session::PendingApproval = + serde_json::from_str(&json).expect("should deserialize without deferred_tool_calls"); + + assert!(parsed.deferred_tool_calls.is_empty()); + assert_eq!(parsed.tool_name, "http"); + assert_eq!(parsed.tool_call_id, "call_123"); + } + + #[test] + fn test_pending_approval_serialization_roundtrip_with_deferred_calls() { + let pending = crate::agent::session::PendingApproval { + request_id: uuid::Uuid::new_v4(), + tool_name: "shell".to_string(), + parameters: serde_json::json!({"command": "echo hi"}), + description: "Run shell command".to_string(), + tool_call_id: "call_1".to_string(), + context_messages: vec![], + deferred_tool_calls: vec![ + ToolCall { + id: "call_2".to_string(), + name: "http".to_string(), + arguments: serde_json::json!({"url": "https://example.com"}), + }, + ToolCall { + id: "call_3".to_string(), + name: "echo".to_string(), + arguments: serde_json::json!({"message": "done"}), + }, + ], + }; + + let json = serde_json::to_string(&pending).expect("serialize"); + let parsed: crate::agent::session::PendingApproval = + serde_json::from_str(&json).expect("deserialize"); + + assert_eq!(parsed.deferred_tool_calls.len(), 2); + assert_eq!(parsed.deferred_tool_calls[0].name, "http"); + assert_eq!(parsed.deferred_tool_calls[1].name, "echo"); + } #[test] fn test_detect_auth_awaiting_positive() { @@ -619,7 +823,7 @@ mod tests { }) .to_string()); - let detected = detect_auth_awaiting("tool_auth", &result); + let detected = check_auth_required("tool_auth", &result); assert!(detected.is_some()); let (name, instructions) = detected.unwrap(); assert_eq!(name, "telegram"); @@ -636,7 +840,7 @@ mod tests { }) .to_string()); - assert!(detect_auth_awaiting("tool_auth", &result).is_none()); + assert!(check_auth_required("tool_auth", &result).is_none()); } #[test] @@ -647,14 +851,14 @@ mod tests { }) .to_string()); - assert!(detect_auth_awaiting("tool_list", &result).is_none()); + assert!(check_auth_required("tool_list", &result).is_none()); } #[test] fn test_detect_auth_awaiting_error_result() { let result: Result = Err(crate::error::ToolError::NotFound { name: "x".into() }.into()); - assert!(detect_auth_awaiting("tool_auth", &result).is_none()); + assert!(check_auth_required("tool_auth", &result).is_none()); } #[test] @@ -666,7 +870,7 @@ mod tests { }) .to_string()); - let (_, instructions) = detect_auth_awaiting("tool_auth", &result).unwrap(); + let (_, instructions) = check_auth_required("tool_auth", &result).unwrap(); assert_eq!(instructions, "Please provide your API token/key."); } @@ -681,7 +885,7 @@ mod tests { }) .to_string()); - let detected = detect_auth_awaiting("tool_activate", &result); + let detected = check_auth_required("tool_activate", &result); assert!(detected.is_some()); let (name, instructions) = detected.unwrap(); assert_eq!(name, "slack"); @@ -697,6 +901,6 @@ mod tests { }) .to_string()); - assert!(detect_auth_awaiting("tool_activate", &result).is_none()); + assert!(check_auth_required("tool_activate", &result).is_none()); } } diff --git a/src/agent/mod.rs b/src/agent/mod.rs index d0c96bc1..1fbbc3bf 100644 --- a/src/agent/mod.rs +++ b/src/agent/mod.rs @@ -44,6 +44,6 @@ pub use self_repair::{BrokenTool, RepairResult, RepairTask, SelfRepair, StuckJob pub use session::{PendingApproval, PendingAuth, Session, Thread, ThreadState, Turn, TurnState}; pub use session_manager::SessionManager; pub use submission::{Submission, SubmissionParser, SubmissionResult}; -pub use task::{Task, TaskContext, TaskHandler, TaskOutput, TaskStatus}; +pub use task::{Task, TaskContext, TaskHandler, TaskOutput}; pub use undo::{Checkpoint, UndoManager}; pub use worker::{Worker, WorkerDeps}; diff --git a/src/agent/routine.rs b/src/agent/routine.rs index 084a9b9f..7fa56d7d 100644 --- a/src/agent/routine.rs +++ b/src/agent/routine.rs @@ -26,6 +26,8 @@ use chrono::{DateTime, Utc}; use serde::{Deserialize, Serialize}; use uuid::Uuid; +use crate::error::RoutineError; + /// A routine is a named, persistent, user-owned task with a trigger and an action. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct Routine { @@ -86,13 +88,16 @@ impl Trigger { } /// Parse a trigger from its DB representation. - pub fn from_db(trigger_type: &str, config: serde_json::Value) -> Result { + pub fn from_db(trigger_type: &str, config: serde_json::Value) -> Result { match trigger_type { "cron" => { let schedule = config .get("schedule") .and_then(|v| v.as_str()) - .ok_or("cron trigger missing 'schedule'")? + .ok_or_else(|| RoutineError::MissingField { + context: "cron trigger".into(), + field: "schedule".into(), + })? .to_string(); Ok(Trigger::Cron { schedule }) } @@ -100,7 +105,10 @@ impl Trigger { let pattern = config .get("pattern") .and_then(|v| v.as_str()) - .ok_or("event trigger missing 'pattern'")? + .ok_or_else(|| RoutineError::MissingField { + context: "event trigger".into(), + field: "pattern".into(), + })? .to_string(); let channel = config .get("channel") @@ -120,7 +128,9 @@ impl Trigger { Ok(Trigger::Webhook { path, secret }) } "manual" => Ok(Trigger::Manual), - other => Err(format!("unknown trigger type: {other}")), + other => Err(RoutineError::UnknownTriggerType { + trigger_type: other.to_string(), + }), } } @@ -186,13 +196,16 @@ impl RoutineAction { } /// Parse an action from its DB representation. - pub fn from_db(action_type: &str, config: serde_json::Value) -> Result { + pub fn from_db(action_type: &str, config: serde_json::Value) -> Result { match action_type { "lightweight" => { let prompt = config .get("prompt") .and_then(|v| v.as_str()) - .ok_or("lightweight action missing 'prompt'")? + .ok_or_else(|| RoutineError::MissingField { + context: "lightweight action".into(), + field: "prompt".into(), + })? .to_string(); let context_paths = config .get("context_paths") @@ -217,12 +230,18 @@ impl RoutineAction { let title = config .get("title") .and_then(|v| v.as_str()) - .ok_or("full_job action missing 'title'")? + .ok_or_else(|| RoutineError::MissingField { + context: "full_job action".into(), + field: "title".into(), + })? .to_string(); let description = config .get("description") .and_then(|v| v.as_str()) - .ok_or("full_job action missing 'description'")? + .ok_or_else(|| RoutineError::MissingField { + context: "full_job action".into(), + field: "description".into(), + })? .to_string(); let max_iterations = config .get("max_iterations") @@ -235,7 +254,9 @@ impl RoutineAction { max_iterations, }) } - other => Err(format!("unknown action type: {other}")), + other => Err(RoutineError::UnknownActionType { + action_type: other.to_string(), + }), } } @@ -334,14 +355,16 @@ impl std::fmt::Display for RunStatus { } impl FromStr for RunStatus { - type Err = String; + type Err = RoutineError; fn from_str(s: &str) -> Result { match s { "running" => Ok(RunStatus::Running), "ok" => Ok(RunStatus::Ok), "attention" => Ok(RunStatus::Attention), "failed" => Ok(RunStatus::Failed), - other => Err(format!("unknown run status: {other}")), + other => Err(RoutineError::UnknownRunStatus { + status: other.to_string(), + }), } } } @@ -370,9 +393,11 @@ pub fn content_hash(content: &str) -> u64 { } /// Parse a cron expression and compute the next fire time from now. -pub fn next_cron_fire(schedule: &str) -> Result>, String> { +pub fn next_cron_fire(schedule: &str) -> Result>, RoutineError> { let cron_schedule = - cron::Schedule::from_str(schedule).map_err(|e| format!("invalid cron: {e}"))?; + cron::Schedule::from_str(schedule).map_err(|e| RoutineError::InvalidCron { + reason: e.to_string(), + })?; Ok(cron_schedule.upcoming(Utc).next()) } diff --git a/src/agent/routine_engine.rs b/src/agent/routine_engine.rs index 52156ac5..93e760f7 100644 --- a/src/agent/routine_engine.rs +++ b/src/agent/routine_engine.rs @@ -25,6 +25,7 @@ use crate::agent::routine::{ use crate::channels::{IncomingMessage, OutgoingResponse}; use crate::config::RoutineConfig; use crate::db::Database; +use crate::error::RoutineError; use crate::llm::{ChatMessage, CompletionRequest, FinishReason, LlmProvider}; use crate::workspace::Workspace; @@ -174,23 +175,26 @@ impl RoutineEngine { } /// Fire a routine manually (from tool call or CLI). - pub async fn fire_manual(&self, routine_id: Uuid) -> Result { + pub async fn fire_manual(&self, routine_id: Uuid) -> Result { let routine = self .store .get_routine(routine_id) .await - .map_err(|e| format!("DB error: {e}"))? - .ok_or_else(|| format!("routine {routine_id} not found"))?; + .map_err(|e| RoutineError::Database { + reason: e.to_string(), + })? + .ok_or(RoutineError::NotFound { id: routine_id })?; if !routine.enabled { - return Err(format!("routine '{}' is disabled", routine.name)); + return Err(RoutineError::Disabled { + name: routine.name.clone(), + }); } if !self.check_concurrent(&routine).await { - return Err(format!( - "routine '{}' already at max concurrent runs", - routine.name - )); + return Err(RoutineError::MaxConcurrent { + name: routine.name.clone(), + }); } let run_id = Uuid::new_v4(); @@ -209,7 +213,9 @@ impl RoutineEngine { }; if let Err(e) = self.store.create_routine_run(&run).await { - return Err(format!("failed to create run record: {e}")); + return Err(RoutineError::Database { + reason: format!("failed to create run record: {e}"), + }); } // Execute inline for manual triggers (caller wants to wait) @@ -313,13 +319,27 @@ async fn execute_routine(ctx: EngineContext, routine: Routine, run: RoutineRun) max_tokens, } => execute_lightweight(&ctx, &routine, prompt, context_paths, *max_tokens).await, RoutineAction::FullJob { description, .. } => { - // Full job mode: for now, execute as lightweight with the description - // as prompt. Full scheduler integration will come as a follow-up. - tracing::info!( + // Full job mode: scheduler integration not yet implemented. + // Execute as lightweight and prepend a warning to the summary. + tracing::warn!( routine = %routine.name, - "FullJob mode executing as lightweight (scheduler integration pending)" + "FullJob mode not yet implemented; falling back to lightweight execution" ); - execute_lightweight(&ctx, &routine, description, &[], ctx.max_lightweight_tokens).await + match execute_lightweight(&ctx, &routine, description, &[], ctx.max_lightweight_tokens) + .await + { + Ok((status, summary, tokens)) => { + let warning = "[Note: FullJob mode is not yet implemented. This routine ran as \ + a single LLM call without tool access. Configure as 'lightweight' \ + or wait for full scheduler integration.]"; + let summary = match summary { + Some(s) => Some(format!("{warning}\n\n{s}")), + None => Some(warning.to_string()), + }; + Ok((status, summary, tokens)) + } + Err(e) => Err(e), + } } }; @@ -331,7 +351,7 @@ async fn execute_routine(ctx: EngineContext, routine: Routine, run: RoutineRun) Ok(execution) => execution, Err(e) => { tracing::error!(routine = %routine.name, "Execution failed: {}", e); - (RunStatus::Failed, Some(e), None) + (RunStatus::Failed, Some(e.to_string()), None) } }; @@ -384,6 +404,20 @@ async fn execute_routine(ctx: EngineContext, routine: Routine, run: RoutineRun) .await; } +/// Sanitize a routine name for use in workspace paths. +/// Only keeps alphanumeric, dash, and underscore characters; replaces everything else. +fn sanitize_routine_name(name: &str) -> String { + name.chars() + .map(|c| { + if c.is_ascii_alphanumeric() || c == '-' || c == '_' { + c + } else { + '_' + } + }) + .collect() +} + /// Execute a lightweight routine (single LLM call). async fn execute_lightweight( ctx: &EngineContext, @@ -391,7 +425,7 @@ async fn execute_lightweight( prompt: &str, context_paths: &[String], max_tokens: u32, -) -> Result<(RunStatus, Option, Option), String> { +) -> Result<(RunStatus, Option, Option), RoutineError> { // Load context from workspace let mut context_parts = Vec::new(); for path in context_paths { @@ -408,8 +442,9 @@ async fn execute_lightweight( } } - // Load routine state from workspace - let state_path = format!("routines/{}/state.md", routine.name); + // Load routine state from workspace (name sanitized to prevent path traversal) + let safe_name = sanitize_routine_name(&routine.name); + let state_path = format!("routines/{safe_name}/state.md"); let state_content = match ctx.workspace.read(&state_path).await { Ok(doc) => Some(doc.content), Err(_) => None, @@ -469,7 +504,9 @@ async fn execute_lightweight( .llm .complete(request) .await - .map_err(|e| format!("LLM call failed: {e}"))?; + .map_err(|e| RoutineError::LlmFailed { + reason: e.to_string(), + })?; let content = response.content.trim(); let tokens_used = Some((response.input_tokens + response.output_tokens) as i32); @@ -477,13 +514,9 @@ async fn execute_lightweight( // Empty content guard (same as heartbeat) if content.is_empty() { return if response.finish_reason == FinishReason::Length { - Err( - "LLM response truncated (finish_reason=length) with no content. \ - Model may have exhausted token budget on reasoning." - .to_string(), - ) + Err(RoutineError::TruncatedResponse) } else { - Err("LLM returned empty content.".to_string()) + Err(RoutineError::EmptyResponse) }; } diff --git a/src/agent/scheduler.rs b/src/agent/scheduler.rs index 23b9ea7c..92b37683 100644 --- a/src/agent/scheduler.rs +++ b/src/agent/scheduler.rs @@ -136,7 +136,9 @@ impl Scheduler { }); // Start the worker - let _ = tx.send(WorkerMessage::Start).await; + if tx.send(WorkerMessage::Start).await.is_err() { + tracing::error!(job_id = %job_id, "Worker died before receiving Start message"); + } // Insert while still holding the write lock jobs.insert(job_id, ScheduledJob { handle, tx }); @@ -418,10 +420,16 @@ impl Scheduler { // Update job state self.context_manager .update_context(job_id, |ctx| { - let _ = ctx.transition_to( + if let Err(e) = ctx.transition_to( JobState::Cancelled, Some("Stopped by scheduler".to_string()), - ); + ) { + tracing::warn!( + job_id = %job_id, + error = %e, + "Failed to transition job to Cancelled state" + ); + } }) .await?; diff --git a/src/agent/self_repair.rs b/src/agent/self_repair.rs index ee7b2a4c..8bb6e19c 100644 --- a/src/agent/self_repair.rs +++ b/src/agent/self_repair.rs @@ -66,12 +66,14 @@ pub trait SelfRepair: Send + Sync { /// Default self-repair implementation. pub struct DefaultSelfRepair { context_manager: Arc, - #[allow(dead_code)] // Will be used for time-based stuck detection + // TODO: use for time-based stuck detection (currently only max_repair_attempts is checked) + #[allow(dead_code)] stuck_threshold: Duration, max_repair_attempts: u32, store: Option>, builder: Option>, - #[allow(dead_code)] // Will be used for tool hot-reload after repair + // TODO: use for tool hot-reload after repair + #[allow(dead_code)] tools: Option>, } @@ -93,15 +95,15 @@ impl DefaultSelfRepair { } /// Add a Store for tool failure tracking. - #[allow(dead_code)] // Public API for configuring repair with persistence - pub fn with_store(mut self, store: Arc) -> Self { + #[allow(dead_code)] // TODO: wire up in main.rs when persistence is needed + pub(crate) fn with_store(mut self, store: Arc) -> Self { self.store = Some(store); self } /// Add a Builder and ToolRegistry for automatic tool repair. - #[allow(dead_code)] // Public API for enabling automatic tool repair - pub fn with_builder( + #[allow(dead_code)] // TODO: wire up in main.rs when auto-repair is needed + pub(crate) fn with_builder( mut self, builder: Arc, tools: Arc, diff --git a/src/agent/session.rs b/src/agent/session.rs index c73882a3..364e6813 100644 --- a/src/agent/session.rs +++ b/src/agent/session.rs @@ -70,10 +70,9 @@ impl Session { pub fn create_thread(&mut self) -> &mut Thread { let thread = Thread::new(self.id); let thread_id = thread.id; - self.threads.insert(thread_id, thread); self.active_thread = Some(thread_id); self.last_active_at = Utc::now(); - self.threads.get_mut(&thread_id).expect("just inserted") + self.threads.entry(thread_id).or_insert(thread) } /// Get the active thread. @@ -88,10 +87,18 @@ impl Session { /// Get or create the active thread. pub fn get_or_create_thread(&mut self) -> &mut Thread { - if self.active_thread.is_none() { - self.create_thread(); + match self.active_thread { + None => self.create_thread(), + Some(id) => { + if self.threads.contains_key(&id) { + self.threads.get_mut(&id).unwrap() + } else { + // Stale active_thread ID: create a new thread, which + // updates self.active_thread to the new thread's ID. + self.create_thread() + } + } } - self.active_thread_mut().expect("just created") } /// Switch to a different thread. @@ -240,7 +247,8 @@ impl Thread { self.turns.push(turn); self.state = ThreadState::Processing; self.updated_at = Utc::now(); - self.turns.last_mut().expect("just pushed") + // turn_number was len() before push, so it's a valid index after push + &mut self.turns[turn_number] } /// Complete the current turn with a response. @@ -353,8 +361,10 @@ impl Thread { if let Some(next) = iter.peek() && next.role == crate::llm::Role::Assistant { - let response = iter.next().expect("peeked"); - turn.complete(&response.content); + // iter.next() is guaranteed Some after a successful peek() + if let Some(response) = iter.next() { + turn.complete(&response.content); + } } self.turns.push(turn); diff --git a/src/agent/session_manager.rs b/src/agent/session_manager.rs index 244348cd..2bce4e8f 100644 --- a/src/agent/session_manager.rs +++ b/src/agent/session_manager.rs @@ -13,6 +13,9 @@ use crate::agent::session::Session; use crate::agent::undo::UndoManager; use crate::hooks::HookRegistry; +/// Warn when session count exceeds this threshold. +const SESSION_COUNT_WARNING_THRESHOLD: usize = 1000; + /// Key for mapping external thread IDs to internal ones. #[derive(Clone, Hash, Eq, PartialEq)] struct ThreadKey { @@ -68,6 +71,14 @@ impl SessionManager { let session = Arc::new(Mutex::new(new_session)); sessions.insert(user_id.to_string(), Arc::clone(&session)); + if sessions.len() >= SESSION_COUNT_WARNING_THRESHOLD && sessions.len() % 100 == 0 { + tracing::warn!( + "High session count: {} active sessions. \ + Pruning runs every 10 minutes; consider reducing session_idle_timeout.", + sessions.len() + ); + } + // Fire OnSessionStart hook (fire-and-forget) if let Some(ref hooks) = self.hooks { let hooks = hooks.clone(); diff --git a/src/agent/submission.rs b/src/agent/submission.rs index de696644..cd1646df 100644 --- a/src/agent/submission.rs +++ b/src/agent/submission.rs @@ -234,6 +234,7 @@ impl Submission { } /// Create an approval submission. + #[cfg(test)] pub fn approval(request_id: Uuid, approved: bool) -> Self { Self::ExecApproval { request_id, @@ -243,6 +244,7 @@ impl Submission { } /// Create an "always approve" submission. + #[cfg(test)] pub fn always_approve(request_id: Uuid) -> Self { Self::ExecApproval { request_id, @@ -252,26 +254,31 @@ impl Submission { } /// Create an interrupt submission. + #[cfg(test)] pub fn interrupt() -> Self { Self::Interrupt } /// Create a compact submission. + #[cfg(test)] pub fn compact() -> Self { Self::Compact } /// Create an undo submission. + #[cfg(test)] pub fn undo() -> Self { Self::Undo } /// Create a redo submission. + #[cfg(test)] pub fn redo() -> Self { Self::Redo } /// Check if this submission starts a new turn. + #[cfg(test)] pub fn starts_turn(&self) -> bool { matches!(self, Self::UserInput { .. }) } @@ -340,6 +347,7 @@ impl SubmissionResult { } /// Create an OK result. + #[cfg(test)] pub fn ok() -> Self { Self::Ok { message: None } } diff --git a/src/agent/task.rs b/src/agent/task.rs index ba5e359c..6d1087c8 100644 --- a/src/agent/task.rs +++ b/src/agent/task.rs @@ -29,6 +29,7 @@ impl TaskOutput { } /// Create a text result. + #[cfg(test)] pub fn text(text: impl Into, duration: Duration) -> Self { Self { result: serde_json::Value::String(text.into()), @@ -37,6 +38,7 @@ impl TaskOutput { } /// Create an empty success result. + #[cfg(test)] pub fn empty(duration: Duration) -> Self { Self { result: serde_json::Value::Null, @@ -130,6 +132,7 @@ impl Task { } /// Create a new Job task with a specific ID. + #[cfg(test)] pub fn job_with_id(id: Uuid, title: impl Into, description: impl Into) -> Self { Self::Job { id, @@ -152,6 +155,7 @@ impl Task { } /// Create a new Background task. + #[cfg(test)] pub fn background(handler: std::sync::Arc) -> Self { Self::Background { id: Uuid::new_v4(), @@ -160,6 +164,7 @@ impl Task { } /// Create a new Background task with a specific ID. + #[cfg(test)] pub fn background_with_id(id: Uuid, handler: std::sync::Arc) -> Self { Self::Background { id, handler } } @@ -174,6 +179,7 @@ impl Task { } /// Get the parent ID for sub-tasks. + #[cfg(test)] pub fn parent_id(&self) -> Option { match self { Self::Job { .. } => None, @@ -225,6 +231,7 @@ impl fmt::Debug for Task { } /// Status of a scheduled task. +#[cfg(test)] #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum TaskStatus { /// Task is queued waiting for execution. diff --git a/src/agent/thread_ops.rs b/src/agent/thread_ops.rs index ba1a6bff..e7db83ca 100644 --- a/src/agent/thread_ops.rs +++ b/src/agent/thread_ops.rs @@ -10,7 +10,7 @@ use uuid::Uuid; use crate::agent::Agent; use crate::agent::compaction::ContextCompactor; -use crate::agent::dispatcher::{AgenticLoopResult, detect_auth_awaiting, parse_auth_result}; +use crate::agent::dispatcher::{AgenticLoopResult, check_auth_required, parse_auth_result}; use crate::agent::session::{PendingApproval, Session, ThreadState}; use crate::agent::submission::SubmissionResult; use crate::channels::{IncomingMessage, StatusUpdate}; @@ -608,8 +608,8 @@ impl Agent { approved: bool, always: bool, ) -> Result { - // Get thread state and pending approval - let (_thread_state, pending) = { + // Get pending approval for this thread + let pending = { let mut sess = session.lock().await; let thread = sess .threads @@ -620,8 +620,7 @@ impl Agent { return Ok(SubmissionResult::error("No pending approval request.")); } - let pending = thread.take_pending_approval(); - (thread.state, pending) + thread.take_pending_approval() }; let pending = match pending { @@ -734,29 +733,17 @@ impl Agent { // If tool_auth returned awaiting_token, enter auth mode and // return instructions directly (skip agentic loop continuation). if let Some((ext_name, instructions)) = - detect_auth_awaiting(&pending.tool_name, &tool_result) + check_auth_required(&pending.tool_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()); - thread.complete_turn(&instructions); - } - } - 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; + self.handle_auth_intercept( + &session, + thread_id, + message, + &tool_result, + ext_name, + instructions.clone(), + ) + .await; return Ok(SubmissionResult::response(instructions)); } @@ -912,29 +899,17 @@ impl Agent { // Auth detection for deferred tools if let Some((ext_name, instructions)) = - detect_auth_awaiting(&tc.name, &deferred_result) + check_auth_required(&tc.name, &deferred_result) { - let auth_data = parse_auth_result(&deferred_result); - { - let mut sess = session.lock().await; - if let Some(thread) = sess.threads.get_mut(&thread_id) { - thread.enter_auth_mode(ext_name.clone()); - thread.complete_turn(&instructions); - } - } - 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; + self.handle_auth_intercept( + &session, + thread_id, + message, + &deferred_result, + ext_name, + instructions.clone(), + ) + .await; return Ok(SubmissionResult::response(instructions)); } @@ -967,7 +942,11 @@ impl Agent { match result { Ok(AgenticLoopResult::Response(response)) => { + let user_input = thread.last_turn().map(|t| t.user_input.clone()); thread.complete_turn(&response); + if let Some(input) = user_input { + self.persist_turn(thread_id, &message.user_id, &input, Some(&response)); + } self.persist_response_chain(thread); let _ = self .channels @@ -1003,16 +982,30 @@ impl Agent { }) } Err(e) => { + let user_input = thread.last_turn().map(|t| t.user_input.clone()); thread.fail_turn(e.to_string()); + if let Some(input) = user_input { + self.persist_turn(thread_id, &message.user_id, &input, None); + } Ok(SubmissionResult::error(e.to_string())) } } } else { - // Rejected - clear approval and return to idle + // Rejected - complete the turn with a rejection message and persist + let rejection = format!( + "Tool '{}' was rejected. The agent will not execute this tool.\n\n\ + You can continue the conversation or try a different approach.", + pending.tool_name + ); { let mut sess = session.lock().await; if let Some(thread) = sess.threads.get_mut(&thread_id) { + let user_input = thread.last_turn().map(|t| t.user_input.clone()); thread.clear_pending_approval(); + thread.complete_turn(&rejection); + if let Some(input) = user_input { + self.persist_turn(thread_id, &message.user_id, &input, Some(&rejection)); + } } } @@ -1025,14 +1018,52 @@ impl Agent { ) .await; - Ok(SubmissionResult::response(format!( - "Tool '{}' was rejected. The agent will not execute this tool.\n\n\ - You can continue the conversation or try a different approach.", - pending.tool_name - ))) + Ok(SubmissionResult::response(rejection)) } } + /// Handle an auth-required result from a tool execution. + /// + /// Enters auth mode on the thread, completes + persists the turn, + /// and sends the AuthRequired status to the channel. + /// Returns the instructions string for the caller to wrap in a response. + async fn handle_auth_intercept( + &self, + session: &Arc>, + thread_id: Uuid, + message: &IncomingMessage, + tool_result: &Result, + ext_name: String, + instructions: String, + ) { + let auth_data = parse_auth_result(tool_result); + { + let mut sess = session.lock().await; + if let Some(thread) = sess.threads.get_mut(&thread_id) { + let user_input = thread.last_turn().map(|t| t.user_input.clone()); + thread.enter_auth_mode(ext_name.clone()); + thread.complete_turn(&instructions); + if let Some(input) = user_input { + self.persist_turn(thread_id, &message.user_id, &input, Some(&instructions)); + } + self.persist_response_chain(thread); + } + } + 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; + } + /// Handle an auth token submitted while the thread is in auth mode. /// /// The token goes directly to the extension manager's credential store, diff --git a/src/agent/undo.rs b/src/agent/undo.rs index 892f7c88..10ab89d5 100644 --- a/src/agent/undo.rs +++ b/src/agent/undo.rs @@ -67,6 +67,7 @@ impl UndoManager { } /// Create with a custom checkpoint limit. + #[cfg(test)] pub fn with_max_checkpoints(mut self, max: usize) -> Self { self.max_checkpoints = max; self @@ -126,6 +127,7 @@ impl UndoManager { } /// Pop the last checkpoint from the undo stack. + #[cfg(test)] pub fn pop_undo(&mut self) -> Option { self.undo_stack.pop_back() } @@ -178,6 +180,7 @@ impl UndoManager { } /// Get a checkpoint by ID. + #[cfg(test)] pub fn get_checkpoint(&self, id: Uuid) -> Option<&Checkpoint> { self.undo_stack .iter() @@ -186,6 +189,7 @@ impl UndoManager { } /// List all available checkpoints (for UI display). + #[cfg(test)] pub fn list_checkpoints(&self) -> Vec<&Checkpoint> { self.undo_stack.iter().collect() } diff --git a/src/agent/worker.rs b/src/agent/worker.rs index 3ec88586..b2ba8a4e 100644 --- a/src/agent/worker.rs +++ b/src/agent/worker.rs @@ -505,7 +505,8 @@ Report when the job is complete or if you encounter issues you cannot resolve."# let output_str = serde_json::to_string_pretty(&output.result) .ok() .map(|s| deps.safety.sanitize_tool_output(tool_name, &s).content); - deps.context_manager + match deps + .context_manager .update_memory(job_id, |mem| { let rec = mem.create_action(tool_name, params.clone()).succeed( output_str.clone(), @@ -516,30 +517,52 @@ Report when the job is complete or if you encounter issues you cannot resolve."# rec }) .await - .ok() + { + Ok(rec) => Some(rec), + Err(e) => { + tracing::warn!(job_id = %job_id, tool = tool_name, "Failed to record action in memory: {e}"); + None + } + } + } + Ok(Err(e)) => { + match deps + .context_manager + .update_memory(job_id, |mem| { + let rec = mem + .create_action(tool_name, params.clone()) + .fail(e.to_string(), elapsed); + mem.record_action(rec.clone()); + rec + }) + .await + { + Ok(rec) => Some(rec), + Err(e) => { + tracing::warn!(job_id = %job_id, tool = tool_name, "Failed to record action in memory: {e}"); + None + } + } + } + Err(_) => { + match deps + .context_manager + .update_memory(job_id, |mem| { + let rec = mem + .create_action(tool_name, params.clone()) + .fail("Execution timeout", elapsed); + mem.record_action(rec.clone()); + rec + }) + .await + { + Ok(rec) => Some(rec), + Err(e) => { + tracing::warn!(job_id = %job_id, tool = tool_name, "Failed to record action in memory: {e}"); + None + } + } } - Ok(Err(e)) => deps - .context_manager - .update_memory(job_id, |mem| { - let rec = mem - .create_action(tool_name, params.clone()) - .fail(e.to_string(), elapsed); - mem.record_action(rec.clone()); - rec - }) - .await - .ok(), - Err(_) => deps - .context_manager - .update_memory(job_id, |mem| { - let rec = mem - .create_action(tool_name, params.clone()) - .fail("Execution timeout", elapsed); - mem.record_action(rec.clone()); - rec - }) - .await - .ok(), }; // Persist action to database (fire-and-forget) diff --git a/src/db/libsql/mod.rs b/src/db/libsql/mod.rs index d83fdbe8..ceae5725 100644 --- a/src/db/libsql/mod.rs +++ b/src/db/libsql/mod.rs @@ -320,10 +320,10 @@ pub(crate) fn row_to_routine_libsql(row: &libsql::Row) -> Result = row.get::(11).ok(); - let trigger = - Trigger::from_db(&trigger_type, trigger_config).map_err(DatabaseError::Serialization)?; + let trigger = Trigger::from_db(&trigger_type, trigger_config) + .map_err(|e| DatabaseError::Serialization(e.to_string()))?; let action = RoutineAction::from_db(&action_type, action_config) - .map_err(DatabaseError::Serialization)?; + .map_err(|e| DatabaseError::Serialization(e.to_string()))?; Ok(Routine { id: get_text(row, 0).parse().unwrap_or_default(), @@ -359,7 +359,7 @@ pub(crate) fn row_to_routine_run_libsql(row: &libsql::Row) -> Result = std::result::Result; diff --git a/src/history/store.rs b/src/history/store.rs index e97ec015..921b4725 100644 --- a/src/history/store.rs +++ b/src/history/store.rs @@ -1179,10 +1179,10 @@ fn row_to_routine(row: &tokio_postgres::Row) -> Result { let max_concurrent: i32 = row.get("max_concurrent"); let dedup_window_secs: Option = row.get("dedup_window_secs"); - let trigger = - Trigger::from_db(&trigger_type, trigger_config).map_err(DatabaseError::Serialization)?; + let trigger = Trigger::from_db(&trigger_type, trigger_config) + .map_err(|e| DatabaseError::Serialization(e.to_string()))?; let action = RoutineAction::from_db(&action_type, action_config) - .map_err(DatabaseError::Serialization)?; + .map_err(|e| DatabaseError::Serialization(e.to_string()))?; Ok(Routine { id: row.get("id"), @@ -1219,7 +1219,7 @@ fn row_to_routine_run(row: &tokio_postgres::Row) -> Result