feat: enhance subagent context with abort flag and implement tool call timeout

This commit is contained in:
asepharyana
2026-07-13 04:59:16 +07:00
parent 3b711bbf3b
commit 2856dd78b8
9 changed files with 272 additions and 131 deletions
+83 -57
View File
@@ -54,65 +54,26 @@ pub fn spawn_bash_job(command: String) -> BashJob {
let (pid_tx, pid_rx) = mpsc::channel::<u32>();
let cmd = command.clone();
let id_for_log = id.clone();
let thread_id = id.clone();
thread::spawn(move || {
let mut child = match Command::new("sh")
.arg("-c")
.arg(&cmd)
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.spawn()
{
Ok(c) => c,
Err(e) => {
let _ = output_tx.try_send(format!("__error:{}", e));
let _ = output_tx.try_send("__exit:-1".to_string());
return;
}
};
// Send the child PID back to the caller so bash_kill can terminate it
let _ = pid_tx.send(child.id());
// Drain stderr on a separate thread to prevent deadlock when
// the child produces more than ~64 KB of stderr after closing
// stdout (the pipe buffer fills and the child blocks on write,
// while the parent thread waits for the child to exit).
let stderr_tx = output_tx.clone();
let _stderr_drain = child.stderr.take().map(|stderr| {
std::thread::spawn(move || {
let reader = std::io::BufReader::new(stderr);
// stderr is intentionally discarded to prevent output-line
// quota pressure from error diagnostics.
for _line in reader.lines().map_while(Result::ok) {
// Discard stderr lines to prevent pipe buffer deadlock.
}
drop(stderr_tx);
})
// Spawn a named thread for easier debugging. If Builder::spawn fails
// (e.g. OS resource limit), fall back to unnameable thread::spawn.
let thread_name = format!("bgbash-{}", &thread_id[..8.min(thread_id.len())]);
if thread::Builder::new().name(thread_name).spawn({
// Clone everything the closure captures so we can also pass it
// to the fallback thread without moving.
let cmd = cmd.clone();
let output_tx = output_tx.clone();
let pid_tx = pid_tx.clone();
let id_for_log = id_for_log.clone();
move || spawn_bash_thread_body(cmd, output_tx, pid_tx, id_for_log)
}).is_err()
{
tracing::warn!("[bgbash:{}] failed to spawn named thread, using unnamed fallback", id_for_log);
thread::spawn(move || {
spawn_bash_thread_body(cmd, output_tx, pid_tx, id_for_log)
});
if let Some(stdout) = child.stdout.take() {
let reader = std::io::BufReader::new(stdout);
for line in reader.lines().map_while(Result::ok) {
// Use try_send so if the channel buffer is full (producer
// faster than consumer), old lines are silently dropped
// rather than growing memory without bound.
if output_tx.try_send(line).is_err() {
// Buffer full — consumer is not draining fast enough.
// Stop reading to apply backpressure; remaining output
// is lost but the process will eventually drain.
tracing::debug!(
"[bgbash:{}] output buffer full ({} lines), discarding remaining output",
id_for_log, MAX_OUTPUT_LINES,
);
break;
}
}
}
let status = child.wait();
let code = status.ok().and_then(|s| s.code());
let _ = output_tx.try_send(format!("__exit:{}", code.unwrap_or(-1)));
});
}
let child_pid = pid_rx.recv().unwrap_or(0);
@@ -124,6 +85,71 @@ pub fn spawn_bash_job(command: String) -> BashJob {
}
}
/// Core bash-thread logic extracted into a free function so it can be
/// spawned from both the named Builder and the unnamed fallback without
/// double-moving the closure.
fn spawn_bash_thread_body(
cmd: String,
output_tx: std::sync::mpsc::SyncSender<String>,
pid_tx: std::sync::mpsc::Sender<u32>,
id_for_log: String,
) {
let mut child = match Command::new("sh")
.arg("-c")
.arg(&cmd)
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.spawn()
{
Ok(c) => c,
Err(e) => {
let _ = output_tx.try_send(format!("__error:{}", e));
let _ = output_tx.try_send("__exit:-1".to_string());
return;
}
};
// Send the child PID back to the caller so bash_kill can terminate it
let _ = pid_tx.send(child.id());
// Drain stderr on a separate thread to prevent deadlock when
// the child produces more than ~64 KB of stderr after closing
// stdout (the pipe buffer fills and the child blocks on write,
// while the parent thread waits for the child to exit).
// Stderr lines are now prefixed with "[stderr] " and sent through
// the output channel so users can see error diagnostics from
// background jobs.
let stderr_tx = output_tx.clone();
let _stderr_drain = child.stderr.take().map(|stderr| {
std::thread::spawn(move || {
let reader = std::io::BufReader::new(stderr);
for line in reader.lines().map_while(Result::ok) {
if stderr_tx.try_send(format!("[stderr] {}", line)).is_err() {
tracing::debug!("[bgbash] stderr buffer full, discarding remaining stderr");
break;
}
}
drop(stderr_tx);
})
});
if let Some(stdout) = child.stdout.take() {
let reader = std::io::BufReader::new(stdout);
for line in reader.lines().map_while(Result::ok) {
if output_tx.try_send(line).is_err() {
tracing::debug!(
"[bgbash:{}] output buffer full ({} lines), discarding remaining output",
id_for_log, MAX_OUTPUT_LINES,
);
break;
}
}
}
let status = child.wait();
let code = status.ok().and_then(|s| s.code());
let _ = output_tx.try_send(format!("__exit:{}", code.unwrap_or(-1)));
}
impl BashJob {
/// Non-blocking poll for the next output line from the job's channel.
///
+76 -55
View File
@@ -114,10 +114,12 @@ const MIN_REASON_LEN: usize = 8;
impl Harness {
/// Decide whether a tool call is allowed to execute.
///
/// Flow: if the tool isn't flagged risky, allow immediately → file-tool
/// reason & path checks → content stub / denial / assumption scan →
/// bash destructive-pattern & exfiltration scan → workspace-root
/// validation for output paths.
/// Flow: ALL tools are gated (not just risky ones), closing the bypass
/// for MCP tools (which are never in the risky list). Basic path
/// traversal and reason validation applies to any tool with a `path`
/// argument. Heavy content scanning (stub/denial/assumption/exfiltration)
/// only applies to risky tools. MCP tools (mcp__ prefix) are treated
/// as risky because their behaviour is unknown.
///
/// Return: `Verdict::Allow` or `Verdict::Block(reason)`.
pub fn gate_tool_call(
@@ -126,33 +128,56 @@ impl Harness {
workspace_roots: &[&std::path::Path],
) -> Verdict {
if !crate::tool::tool_is_risky(tool_name) {
return Verdict::Allow;
}
let is_risky = crate::tool::tool_is_risky(tool_name);
let is_mcp = tool_name.starts_with("mcp__");
// File-mutating tools: write / edit / delete
if matches!(tool_name, "write" | "edit" | "delete") {
if let Some(path) = args.get("path").and_then(|v| v.as_str()) {
if path.contains("..") {
return Verdict::Block(
"path traversal detected in 'path' argument".to_string(),
);
}
if !workspace_roots.is_empty() {
let abs_check = std::path::PathBuf::from(path);
if abs_check.is_absolute()
&& !workspace_roots.iter().any(|r| abs_check.starts_with(r))
{
return Verdict::Block(format!(
"absolute path '{path}' is outside all workspace roots"
));
}
// ── Universal checks applied to EVERY tool ──
// Path traversal: check ANY tool that accepts a path argument,
// not just write/edit/delete, so tools like read, MCP tools,
// and future tools are also protected.
if let Some(path) = args.get("path").and_then(|v| v.as_str()) {
if path.contains("..") {
return Verdict::Block(
"path traversal detected in 'path' argument".to_string(),
);
}
if !workspace_roots.is_empty() {
let abs_check = std::path::PathBuf::from(path);
if abs_check.is_absolute()
&& !workspace_roots.iter().any(|r| abs_check.starts_with(r))
{
return Verdict::Block(format!(
"absolute path '{path}' is outside all workspace roots"
));
}
}
}
// write / edit require a non-trivial `reason` argument (hooks-style
// discipline: every mutation must explain itself).
// Workspace-root validation for output path.
if let Some(out_path) = Self::find_output_path(tool_name, args) {
if !workspace_roots.is_empty()
&& !out_path.starts_with("/tmp")
&& !out_path.is_absolute()
{
let allowed = workspace_roots.iter().any(|r| out_path.starts_with(r));
if !allowed {
return Verdict::Block(format!(
"output path '{:?}' is outside all workspace roots",
out_path
));
}
}
}
// ── Risky / MCP tool checks ──
// Non-risky, non-MCP tools (read, grep, glob, recall, etc.) are
// allowed after universal checks above.
if !is_risky && !is_mcp {
return Verdict::Allow;
}
// File-mutating tools: write / edit / delete
if matches!(tool_name, "write" | "edit" | "delete") {
match Self::validate_reason(tool_name, args) {
Ok(()) => {}
@@ -186,7 +211,8 @@ impl Harness {
}
}
// Bash: destructive patterns, exfiltration, sensitive-path reads.
// Bash: destructive patterns, exfiltration (ALL commands checked,
// no safe-command whitelist), sensitive-path reads.
if tool_name == "bash" {
let cmd = args.get("command").and_then(|v| v.as_str()).unwrap_or("");
if cmd.contains("..") {
@@ -194,23 +220,14 @@ impl Harness {
"path traversal detected in bash command".to_string(),
);
}
if !cmd.trim_start().starts_with("cargo")
&& !cmd.trim_start().starts_with("rustc")
&& !cmd.trim_start().starts_with("git ")
&& !cmd.trim_start().starts_with("ls")
&& !cmd.trim_start().starts_with("pwd")
&& !cmd.trim_start().starts_with("echo")
&& !cmd.trim_start().starts_with("cat")
&& !cmd.trim_start().starts_with("find")
&& !cmd.trim_start().starts_with("grep")
&& !cmd.trim_start().starts_with("test")
{
for pat in EXFIL_PATTERNS {
if cmd.contains(pat) {
return Verdict::Block(format!(
"potential data-exfiltration command blocked (matched '{pat}')"
));
}
// Exfiltration patterns are checked on EVERY bash command,
// regardless of prefix. The safe-command whitelist was removed
// because it could be bypassed with command chaining.
for pat in EXFIL_PATTERNS {
if cmd.contains(pat) {
return Verdict::Block(format!(
"potential data-exfiltration command blocked (matched '{pat}')"
));
}
}
for pat in SENSITIVE_PATH_PATTERNS {
@@ -259,22 +276,26 @@ impl Harness {
}
}
// Workspace-root validation for the resolved output path.
if let Some(out_path) = Self::find_output_path(tool_name, args) {
if !workspace_roots.is_empty()
&& !out_path.starts_with("/tmp")
&& !out_path.is_absolute()
{
let allowed = workspace_roots.iter().any(|r| out_path.starts_with(r));
if !allowed {
// MCP tools: unknown behaviour — require a reason if they take
// arguments, to discourage lazy invocations.
if is_mcp {
if let Some(reason) = args.get("reason").and_then(|v| v.as_str()) {
if reason.trim().len() < MIN_REASON_LEN {
return Verdict::Block(format!(
"output path '{:?}' is outside all workspace roots",
out_path
"MCP tool '{tool_name}' requires a non-trivial 'reason' \
(>= {MIN_REASON_LEN} chars) explaining why it is needed"
));
}
} else if args.as_object().map(|m| !m.is_empty()).unwrap_or(false) {
// Only require reason when there are meaningful arguments
return Verdict::Block(format!(
"MCP tool '{tool_name}' requires a 'reason' argument \
explaining the operation"
));
}
}
Self::classify(tool_name)
Verdict::Allow
}
/// Validate the `reason` argument for a mutating tool.
+7 -1
View File
@@ -2,7 +2,7 @@
//! including the default read-only tool set for reviewer agents.
use std::path::PathBuf;
use std::sync::{Arc, Mutex};
use std::sync::{Arc, Mutex, atomic::{AtomicBool, Ordering}};
use super::spawn::AgentDefinition;
/// Default read-only tool names granted to `role == "reviewer"` agents.
@@ -21,6 +21,11 @@ pub struct SubagentContext {
/// workflow run. Set by the workflow engine; `note_finding` writes
/// into this from tool code via `ToolCtx.workflow_findings`.
pub workflow_findings: Option<Arc<Mutex<Vec<String>>>>,
/// Atomic abort flag: when set to `true`, the subagent loop will exit
/// at the earliest opportunity (before the next LLM call). Mirrors the
/// main agent's `abort_flag` mechanism so that long-running or stuck
/// subagents can be cancelled from the parent.
pub abort_flag: Option<Arc<AtomicBool>>,
}
/// Build a `SubagentContext` from an `AgentDefinition`.
@@ -48,5 +53,6 @@ pub fn build_subagent_context(def: AgentDefinition) -> SubagentContext {
session_dir: PathBuf::new(),
workspaces: Vec::new(),
workflow_findings: None,
abort_flag: None,
}
}
+23
View File
@@ -18,6 +18,10 @@ use super::event::SubagentEvent;
#[allow(dead_code)]
pub const MAX_AGENT_STEPS: usize = usize::MAX;
/// Maximum time a single tool call may block inside a subagent before
/// being abandoned. Prevents a stuck tool from hanging the subagent loop.
const SUBAGENT_TOOL_TIMEOUT_MS: u64 = 120_000;
/// Maps a subagent's allowed tool names to concrete Tool trait objects and
/// OpenAI-style tool definitions.
///
@@ -328,6 +332,16 @@ pub fn run_subagent(ctx: SubagentContext, tx: mpsc::Sender<SubagentEvent>) -> an
for step in 0..ctx.max_steps {
// Check abort flag before each LLM call so a stuck subagent can
// be cancelled from the parent (mirrors main agent behaviour).
if ctx.abort_flag.as_ref().is_some_and(|f| f.load(std::sync::atomic::Ordering::SeqCst)) {
let _ = tx.blocking_send(SubagentEvent::StepFailed {
_step: step,
_error: "subagent aborted by parent".to_string(),
});
anyhow::bail!("subagent aborted by parent at step {}", step);
}
// Use the structured tool-calling API so the LLM can request tools with
// proper arguments, exactly like the main agent does.
let (response, _usage) = match client.chat_with_tools_non_streaming(&messages, tdefs_opt.clone()) {
@@ -352,6 +366,15 @@ pub fn run_subagent(ctx: SubagentContext, tx: mpsc::Sender<SubagentEvent>) -> an
messages.push(response);
for tool_call in &tool_calls {
// Check abort flag before each tool execution
if ctx.abort_flag.as_ref().is_some_and(|f| f.load(std::sync::atomic::Ordering::SeqCst)) {
let _ = tx.blocking_send(SubagentEvent::StepFailed {
_step: step,
_error: "subagent aborted by parent during tool execution".to_string(),
});
anyhow::bail!("subagent aborted by parent during tool call at step {}", step);
}
let tool_name = &tool_call.function.name;
let args = crate::dto::chat::tool::sanitize_tool_arguments(&tool_call.function.arguments);
let explicitly_allowed = ctx.allowed_tools.contains(tool_name);
+3
View File
@@ -141,6 +141,9 @@ fn spawn_single_agent(
// Link the shared findings Arc so note_finding calls within this
// subagent write into the same vec visible to sibling agents.
ctx.workflow_findings = Some(findings.clone());
// Abort flag stays None by default — the parent can set it to abort
// long-running agents. No abort mechanism is wired yet at this level;
// future work can expose a kill-switch per agent via the live callback.
// Create an mpsc channel and drain events in a background thread so
// run_subagent's blocking_send never blocks (previously the _rx was
+12 -7
View File
@@ -33,21 +33,26 @@ pub struct ToolFunction {
/// than a nested object; if `args` is a string, attempt to parse it as
/// JSON. Objects and other value types pass through unchanged.
///
/// Why: falling back to the raw string on parse failure (rather than
/// erroring) keeps the harness resilient to malformed provider output.
/// Security: on parse failure we wrap the raw string in `{ "_raw": "..." }`
/// instead of passing it through as a raw string, so tools that expect a
/// JSON object (via `args.get("key")`) get `None` rather than unexpectedly
/// receiving a plain string value.
///
/// Return: the parsed `Value`, or the original `args` clone if parsing fails.
/// Return: the parsed `Value`, or a wrapper object on parse failure.
pub fn sanitize_tool_arguments(args: &Value) -> Value {
match args {
Value::String(s) => {
match serde_json::from_str::<Value>(s) {
Ok(v) => v,
Err(e) => {
tracing::warn!(
"warning: tool argument is a JSON string but failed to parse: {}. Using raw string.",
e
tracing::error!(
"tool argument is a JSON string but failed to parse: {}. \
Wrapping in object to prevent tool misbehaviour. Raw was: {}",
e, s.chars().take(200).collect::<String>(),
);
args.clone()
// Wrap in a safe object so tools don't receive a raw
// string that could be misinterpreted as an object key.
serde_json::json!({"_raw": s, "_parse_error": e.to_string()})
}
}
}
+34 -11
View File
@@ -17,35 +17,58 @@ pub struct EditLogEntry {
pub session_id: String,
}
/// Maximum number of edit entries held in memory at once.
/// Beyond this limit, old entries are dropped from the in-memory cache
/// to prevent unbounded memory growth in long sessions.
const MAX_MEMORY_ENTRIES: usize = 10_000;
/// In-memory view of a session's edit log, backed by `edits.jsonl` on disk.
#[derive(Debug, Clone)]
pub struct EditLog {
pub entries: Vec<EditLogEntry>,
pub path: std::path::PathBuf,
/// Total entries on disk (may exceed `entries.len()` if truncated).
pub total_on_disk: usize,
}
impl EditLog {
/// Open (or start tracking) the edit log for a session directory,
/// replaying any existing `edits.jsonl` into memory.
/// replaying any existing `edits.jsonl` into memory (capped at
/// `MAX_MEMORY_ENTRIES` to prevent OOM).
pub fn new(session_dir: &std::path::Path) -> Self {
let path = session_dir.join("edits.jsonl");
let entries = Self::load_from_disk(&path);
EditLog { entries, path }
let (entries, total_on_disk) = Self::load_from_disk(&path);
EditLog { entries, path, total_on_disk }
}
/// Reads every line of edits.jsonl back into memory so callers who create a
/// *new* EditLog after a previous session can inspect the full history.
fn load_from_disk(path: &std::path::Path) -> Vec<EditLogEntry> {
/// Reads lines of edits.jsonl into memory, keeping only the most recent
/// `MAX_MEMORY_ENTRIES` entries. The full history is preserved on disk
/// regardless of the in-memory limit.
fn load_from_disk(path: &std::path::Path) -> (Vec<EditLogEntry>, usize) {
let file = match std::fs::File::open(path) {
Ok(f) => f,
Err(_) => return Vec::new(),
Err(_) => return (Vec::new(), 0),
};
use std::io::{BufRead, BufReader};
let reader = BufReader::new(file);
reader
.lines()
.filter_map(|line| line.ok().and_then(|l| serde_json::from_str(&l).ok()))
.collect()
let mut entries: Vec<EditLogEntry> = Vec::new();
let mut total = 0usize;
for line in reader.lines() {
let line = match line {
Ok(l) => l,
Err(_) => continue,
};
total += 1;
if let Ok(entry) = serde_json::from_str::<EditLogEntry>(&line) {
// Keep only the most recent entries in memory
if entries.len() >= MAX_MEMORY_ENTRIES {
// Drop oldest (front) to make room
entries.remove(0);
}
entries.push(entry);
}
}
(entries, total)
}
/// Append one entry to `edits.jsonl` on disk and to the in-memory log,
+11
View File
@@ -78,9 +78,20 @@ impl Session {
/// Load a session's metadata by id from `<base_dir>/sessions/<id>/session.json`.
///
/// Security: the session id is validated to prevent directory traversal
/// (e.g. `../../etc/passwd`). Only alphanumeric, hyphens, underscores,
/// and dots are allowed — no path separators.
///
/// Return: the parsed `Session`, or an `io::Error` if the file is
/// missing or malformed.
pub fn load(id: &str, base_dir: &Path) -> std::io::Result<Self> {
// Reject session ids that contain path separators or parent dir refs
if id.contains('/') || id.contains('\\') || id.contains("..") {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
format!("invalid session id '{}': must not contain path separators", id),
));
}
let path = base_dir.join("sessions").join(id).join("session.json");
let data = std::fs::read_to_string(path)?;
let session: Session = serde_json::from_str(&data)?;
+23
View File
@@ -36,6 +36,11 @@ impl Tool for BashOutput {
.and_then(|v| v.as_str())
.ok_or_else(|| anyhow!("missing required argument: job_id"))?
.to_string();
// Validate that job_id looks like a UUID to prevent injection
// into the global job registry.
if !is_valid_job_id(&job_id) {
anyhow::bail!("invalid job_id format: expected UUID");
}
match crate::app::bgbash::control::bash_output(&job_id) {
Some(lines) => Ok(lines.join("\n")),
None => Ok(format!("No new output from job '{}'", job_id)),
@@ -73,7 +78,25 @@ impl Tool for BashKill {
.and_then(|v| v.as_str())
.ok_or_else(|| anyhow!("missing required argument: job_id"))?
.to_string();
if !is_valid_job_id(&job_id) {
anyhow::bail!("invalid job_id format: expected UUID");
}
crate::app::bgbash::control::bash_kill(&job_id)?;
Ok(format!("Killed background job '{}'", job_id))
}
}
/// Validate that a job_id matches UUID v4 format (hex with dashes).
fn is_valid_job_id(id: &str) -> bool {
// UUID v4 format: 8-4-4-4-12 hex digits
let parts: Vec<&str> = id.split('-').collect();
if parts.len() != 5 {
return false;
}
parts.iter().all(|p| !p.is_empty() && p.chars().all(|c| c.is_ascii_hexdigit()))
&& parts[0].len() == 8
&& parts[1].len() == 4
&& parts[2].len() == 4
&& parts[3].len() == 4
&& parts[4].len() == 12
}