Compare commits

..
Author SHA1 Message Date
Claude a5b5d02ab1 fix(clippy): replace find().is_none() with !any() in user stats test
https://claude.ai/code/session_01Nm95eCjdrwxDjwkHTRieZs
2026-03-28 19:22:38 +00:00
ZakiandClaude a0020b22a5 fix(routines): address review feedback on retry loop
- Remove outer retry for LlmFailed errors since RetryProvider already
  handles transient LLM failures with its own bounded budget, preventing
  multiplicative retry counts

- Preserve Option<i32> semantics for token accumulation: None means
  "unknown/not tracked" rather than converting to Some(0) via
  unwrap_or(0), so downstream API/UI correctly distinguishes null
  from zero

- Add partial_tokens field to EmptyResponse and TruncatedResponse
  variants so token usage from those failed attempts is captured in the
  retry accumulator

- Persist accumulated token total on final failure path so usage from
  earlier retry attempts is not silently discarded

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-28 19:16:10 +00:00
Claude b3e09c3827 style: fix rustfmt formatting in routine_engine.rs
https://claude.ai/code/session_01CsP5wMZ2evEMHgGghjAfR1
2026-03-28 19:16:10 +00:00
ZakiandClaude eba088f30e fix(routines): address PR review feedback for bounded retry (#1320)
- Skip retry for tools-enabled routines to prevent duplicate side effects
- Replace fragile PERMANENT_LLM_PATTERNS substring matching with a
  retryable bool field on RoutineError::LlmFailed, set at the LlmError
  conversion site using llm::retry::is_retryable()
- Use saturating_add for token accumulation

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-28 19:16:10 +00:00
ZakiandClaude 0d82ce5d4c fix: use structured tracing for transient routine retry logging [skip-regression-check]
Replace tracing::warn! with tracing::event! targeting "transient_routine_errors"
for better log filtering and structured field capture on retry attempts.

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-28 19:16:10 +00:00
ZakiandClaude 3a27cd3561 fix(routines): classify permanent LLM failures, accumulate tokens, warn on tool retry (#1320)
Address three review comments on the bounded-retry implementation:

1. Permanent LLM failures no longer retried: `is_retryable()` now inspects
   the `LlmFailed` reason string for known permanent patterns (auth,
   content policy, context length, model not available, moderation).
   These map to `LlmError` variants already classified as non-retryable
   by the LLM retry layer.

2. Token usage accumulated across retries: `RoutineError::LlmFailed` now
   carries an optional `partial_tokens` field populated from tokens consumed
   before the failure. The retry loop sums partial tokens from failed
   attempts with the final successful attempt's tokens.

3. Tool loop retry limitation documented: added a code comment explaining
   that the retry wraps the entire `execute_lightweight()` call, and a
   warning log when retrying a tools-enabled routine so operators know
   side effects may be repeated.

[skip-regression-check]

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-28 19:16:10 +00:00
Claude c20fd6ec3a fix(deps): resolve RUSTSEC-2026-0049 rustls-webpki CRL advisory
Update rustls-webpki 0.103.9 -> 0.103.10 and ignore the advisory for
0.102.8 which is pinned by libsql's rustls 0.22.4 dependency chain.

https://claude.ai/code/session_01MmxvBgAMn4m45pZguFKBEX
2026-03-28 19:16:09 +00:00
Claude 18026fcb7c style: run cargo fmt to fix formatting
https://claude.ai/code/session_01Nv2TJ3so5WQqpRUcirAhT3
2026-03-28 19:16:09 +00:00
ZakiandClaude 841cf5fe76 fix(routines): add bounded retry for transient lightweight execution failures (#1320)
Lightweight routine executions that fail with transient errors (LLM
failures, empty responses, truncated responses) are now retried up to 3
times with exponential backoff (1s, 2s, 4s) before reporting failure.

Full-job routines are not retried since the scheduler/watcher handles
their lifecycle. Hard failures (disabled, not found, auth, DB errors)
fail immediately without retry.

Adds RoutineError::is_retryable() to classify transient vs hard errors,
with comprehensive regression tests covering all error variants.

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-28 19:16:09 +00:00
fd41bdf4be fix(worker): treat empty LLM response after text output as completion (#1677)
* fix(worker): treat empty LLM response after text output as completion

When a job's LLM produces a substantive text response (e.g., formatted
results from a routine) and the next LLM call returns empty or errors,
the worker now treats this as successful completion instead of
continuing the loop until failure.

Previously, empty responses always triggered TextAction::Continue,
causing the loop to re-call the LLM. The LLM had nothing more to say,
so the provider returned "Response contained no message or tool call
(empty)". This made routine jobs that successfully produced results
report as "failed".

The fix adds a `has_text_response` flag to JobDelegate:
- After any non-empty text response: flag is set
- Empty text after flag is set: treated as completion
- LLM errors (select_tools/respond_with_tools) after flag: treated
  as completion instead of propagating
- Empty text before any output: still retries (rate-limit backoff)

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>

* fix(worker): restrict error swallowing to EmptyResponse variant only

- Add LlmError::EmptyResponse variant for when LLM returns no content
- Update nearai_chat and github_copilot providers to emit EmptyResponse
  instead of InvalidResponse for empty/no-choice responses
- try_complete_on_error now only swallows EmptyResponse (not AuthFailed,
  ContextLengthExceeded, Http, Io, etc.)
- Extract is_completion_eligible_error as testable pure function
- Log mark_completed errors at warn level instead of silently dropping
- Add EmptyResponse to retry and circuit breaker transient classifications
- Rewrite test to exercise real classification logic against all variants

Addresses review feedback from zmanian and gemini-code-assist.

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>

* refactor(worker): extract mark_completed_or_warn helper to DRY completion logic

Extract shared mark-completed + warn-on-failure pattern into a single
helper method used by both try_complete_on_error and handle_text_response.

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>

---------

Co-authored-by: j-bloggs <[email protected]>
Co-authored-by: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-28 18:46:08 +01:00
de5a1c7b0d fix(worker): replace script -qfc with pty-process for injection-safe PTY (#1678)
- Add pty-process crate (MIT, tokio async support) for PTY allocation
- Spawn claude CLI with pty-process::Command::arg() chaining instead of
  building a shell string for script -qfc
- Eliminates all shell injection surfaces: prompt, model, session_id
  are passed via execve, never interpreted by a shell
- Keep stderr on separate pipe to prevent NDJSON parse breakage
  (pty-process attaches PTY to all fds by default)
- Gate PTY behind #[cfg(unix)] with direct-spawn fallback for Windows CI
- Read stdout from PTY master (implements tokio::io::AsyncRead)
- Add regression tests: arg vector construction + PTY allocation

Addresses review feedback from zmanian and gemini-code-assist.

Co-authored-by: j-bloggs <[email protected]>
Co-authored-by: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-28 16:31:49 +01:00
9ce3a9fc53 feat(discord): implement on_broadcast via DM channel creation (#1693)
- Implement broadcast_dm() that creates a DM channel with the target
  user (POST /users/@me/channels, cached by Discord) and sends the
  message to it
- Extract DISCORD_API_BASE constant for all Discord REST API URLs
- Extract send_channel_message() shared helper to deduplicate message
  posting between on_respond and broadcast_dm
- Add snowflake validation on user_id before API calls
- Fix pre-existing clippy redundant_closure warning
- Use typed DmChannelResponse struct instead of serde_json::Value

Closes no specific issue — completes the previously stubbed on_broadcast.

Co-authored-by: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-28 16:31:27 +01:00
31 changed files with 913 additions and 525 deletions
Generated
+16 -5
View File
@@ -3150,7 +3150,7 @@ dependencies = [
"libc",
"percent-encoding",
"pin-project-lite",
"socket2 0.5.10",
"socket2 0.6.3",
"system-configuration",
"tokio",
"tower-service",
@@ -3439,6 +3439,7 @@ dependencies = [
"pgvector",
"postgres-types",
"pretty_assertions",
"pty-process",
"rand 0.8.5",
"readabilityrs",
"refinery",
@@ -3524,7 +3525,7 @@ checksum = "3640c1c38b8e4e43584d8df18be5fc6b0aa314ce6ebf51b53313d4306cca8e46"
dependencies = [
"hermit-abi",
"libc",
"windows-sys 0.59.0",
"windows-sys 0.61.2",
]
[[package]]
@@ -4906,6 +4907,16 @@ dependencies = [
"syn 1.0.109",
]
[[package]]
name = "pty-process"
version = "0.5.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "71cec9e2670207c5ebb9e477763c74436af3b9091dd550b9fb3c1bec7f3ea266"
dependencies = [
"rustix 1.1.4",
"tokio",
]
[[package]]
name = "pulley-interpreter"
version = "28.0.1"
@@ -4930,7 +4941,7 @@ dependencies = [
"quinn-udp",
"rustc-hash 2.1.1",
"rustls 0.23.37",
"socket2 0.5.10",
"socket2 0.6.3",
"thiserror 2.0.18",
"tokio",
"tracing",
@@ -4967,9 +4978,9 @@ dependencies = [
"cfg_aliases",
"libc",
"once_cell",
"socket2 0.5.10",
"socket2 0.6.3",
"tracing",
"windows-sys 0.59.0",
"windows-sys 0.60.2",
]
[[package]]
+4
View File
@@ -189,6 +189,10 @@ json5 = { version = "0.4", optional = true }
[target.'cfg(target_os = "macos")'.dependencies]
security-framework = "3"
# PTY allocation for Claude CLI stdout buffering fix (Unix only)
[target.'cfg(unix)'.dependencies]
pty-process = { version = "0.5", features = ["async"] }
# Linux secret-service (GNOME Keyring, KWallet)
[target.'cfg(target_os = "linux")'.dependencies]
secret-service = { version = "4", features = ["rt-tokio-crypto-rust"] }
+118 -23
View File
@@ -28,6 +28,9 @@ use std::{cmp::Ordering, collections::HashMap};
use ed25519_dalek::{Signature, Verifier, VerifyingKey};
use serde::{Deserialize, Serialize};
/// Discord REST API v10 base URL.
const DISCORD_API_BASE: &str = "https://discord.com/api/v10";
use exports::near::agent::channel::{
AgentResponse, ChannelConfig, Guest, HttpEndpointConfig, IncomingHttpRequest,
OutgoingHttpResponse, PollConfig, StatusUpdate,
@@ -427,7 +430,7 @@ impl Guest for DiscordChannel {
(
"PATCH",
format!(
"https://discord.com/api/v10/webhooks/{}/{}/messages/@original",
"{DISCORD_API_BASE}/webhooks/{}/{}/messages/@original",
application_id, token
),
)
@@ -438,20 +441,7 @@ impl Guest for DiscordChannel {
payload["allowed_mentions"] = serde_json::json!({
"replied_user": true
});
let mention_payload = serde_json::to_vec(&payload)
.map_err(|e| format!("Failed to serialize mention payload: {}", e))?;
let mention_url = format!(
"https://discord.com/api/v10/channels/{}/messages",
metadata.channel_id
);
let result = channel_host::http_request(
"POST",
&mention_url,
&discord_auth_headers_json(true),
Some(&mention_payload),
None,
);
return map_discord_response(result);
return send_channel_message(&metadata.channel_id, payload);
} else {
return Err("Unsupported Discord response metadata".to_string());
};
@@ -469,8 +459,8 @@ impl Guest for DiscordChannel {
fn on_status(_update: StatusUpdate) {}
fn on_broadcast(_user_id: String, _response: AgentResponse) -> Result<(), String> {
Err("broadcast not yet implemented for Discord channel".to_string())
fn on_broadcast(user_id: String, response: AgentResponse) -> Result<(), String> {
broadcast_dm(&user_id, &response.content)
}
fn on_shutdown() {
@@ -501,6 +491,21 @@ fn map_discord_response(
}
}
/// Post a JSON payload to a Discord channel as a new message.
fn send_channel_message(channel_id: &str, payload: serde_json::Value) -> Result<(), String> {
let payload_bytes = serde_json::to_vec(&payload)
.map_err(|e| format!("Failed to serialize message: {}", e))?;
let url = format!("{DISCORD_API_BASE}/channels/{}/messages", channel_id);
let result = channel_host::http_request(
"POST",
&url,
&discord_auth_headers_json(true),
Some(&payload_bytes),
None,
);
map_discord_response(result)
}
fn load_runtime_config() -> DiscordRuntimeConfig {
channel_host::workspace_read("config.json")
.and_then(|raw| serde_json::from_str::<DiscordRuntimeConfig>(&raw).ok())
@@ -539,7 +544,7 @@ fn get_or_fetch_bot_id() -> Option<String> {
let response = channel_host::http_request(
"GET",
"https://discord.com/api/v10/users/@me",
&format!("{DISCORD_API_BASE}/users/@me"),
&discord_auth_headers_json(false),
None,
Some(10_000),
@@ -659,7 +664,7 @@ fn poll_channel_mentions(channel_id: &str, bot_id: &str) {
fn fetch_latest_message_id(channel_id: &str) -> Option<String> {
let url = format!(
"https://discord.com/api/v10/channels/{}/messages?limit=1",
"{DISCORD_API_BASE}/channels/{}/messages?limit=1",
channel_id
);
let response = channel_host::http_request(
@@ -697,7 +702,7 @@ fn fetch_messages_after_cursor(
for page in 0..MAX_PAGES {
let url = format!(
"https://discord.com/api/v10/channels/{}/messages?limit={}&after={}",
"{DISCORD_API_BASE}/channels/{}/messages?limit={}&after={}",
channel_id, PAGE_LIMIT, after
);
let response = match channel_host::http_request(
@@ -986,7 +991,7 @@ fn handle_slash_command(interaction: &DiscordInteraction) -> bool {
);
// Attempt to notify user of internal error
let url = format!(
"https://discord.com/api/v10/webhooks/{}/{}",
"{DISCORD_API_BASE}/webhooks/{}/{}",
interaction.application_id, interaction.token
);
let payload = serde_json::json!({
@@ -1106,7 +1111,7 @@ fn check_sender_permission(
}
let dm_policy =
channel_host::workspace_read(DM_POLICY_PATH).unwrap_or_else(|| default_dm_policy());
channel_host::workspace_read(DM_POLICY_PATH).unwrap_or_else(default_dm_policy);
if dm_policy == "open" {
return true;
}
@@ -1161,7 +1166,7 @@ fn check_sender_permission(
/// Send a pairing code as an ephemeral Discord followup message.
fn send_pairing_reply(ctx: &PairingReplyCtx, code: &str) -> Result<(), String> {
let url = format!(
"https://discord.com/api/v10/webhooks/{}/{}",
"{DISCORD_API_BASE}/webhooks/{}/{}",
ctx.application_id, ctx.token
);
let payload = serde_json::json!({
@@ -1194,6 +1199,57 @@ fn send_pairing_reply(ctx: &PairingReplyCtx, code: &str) -> Result<(), String> {
}
}
/// Send a broadcast message to a Discord user via DM.
///
/// Creates a DM channel with the user (Discord caches this, so repeated calls
/// for the same user reuse the existing channel) and then posts the message.
fn broadcast_dm(user_id: &str, content: &str) -> Result<(), String> {
// Validate user_id is a plausible Discord snowflake (numeric, 17-20 digits)
// to avoid injecting arbitrary strings into API URLs.
if user_id.is_empty()
|| !user_id.chars().all(|c| c.is_ascii_digit())
|| user_id.len() < 17
|| user_id.len() > 20
{
return Err(format!("Invalid Discord user ID: '{}'", user_id));
}
// Step 1: Open (or reuse) a DM channel with the target user.
let create_dm_payload = serde_json::json!({ "recipient_id": user_id });
let create_dm_bytes = serde_json::to_vec(&create_dm_payload)
.map_err(|e| format!("Failed to serialize DM channel request: {}", e))?;
let dm_response = channel_host::http_request(
"POST",
&format!("{DISCORD_API_BASE}/users/@me/channels"),
&discord_auth_headers_json(true),
Some(&create_dm_bytes),
Some(10_000),
)
.map_err(|e| format!("Failed to create DM channel: {}", e))?;
if !(200..300).contains(&dm_response.status) {
let body = String::from_utf8_lossy(&dm_response.body);
return Err(format!(
"Discord create-DM failed: {} - {}",
dm_response.status, body
));
}
#[derive(Deserialize)]
struct DmChannelResponse {
id: String,
}
let dm_channel: DmChannelResponse = serde_json::from_slice(&dm_response.body)
.map_err(|e| format!("Failed to parse DM channel response: {}", e))?;
let channel_id = &dm_channel.id;
// Step 2: Send the message to the DM channel.
let truncated = truncate_message(content);
let payload = serde_json::json!({ "content": truncated });
send_channel_message(channel_id, payload)
}
fn json_response(status: u16, value: serde_json::Value) -> OutgoingHttpResponse {
let body = serde_json::to_vec(&value).unwrap_or_default();
let headers = serde_json::json!({"Content-Type": "application/json"});
@@ -1593,4 +1649,43 @@ mod tests {
assert_eq!(interaction.interaction_type, 2);
assert!(interaction.data.is_some());
}
#[test]
fn test_broadcast_dm_payload_format() {
// Verify the DM channel creation payload is well-formed JSON that
// Discord's API expects.
let user_id = "123456789012345678";
let payload = serde_json::json!({ "recipient_id": user_id });
let serialized = serde_json::to_vec(&payload).unwrap();
let parsed: serde_json::Value = serde_json::from_slice(&serialized).unwrap();
assert_eq!(
parsed.get("recipient_id").and_then(|v| v.as_str()),
Some(user_id)
);
}
#[test]
fn test_broadcast_message_truncation() {
// Broadcast uses truncate_message, verify it handles content within
// Discord's 2000-char limit for DMs.
let short = "Hello from broadcast";
assert_eq!(truncate_message(short), short);
let long = "x".repeat(2500);
let result = truncate_message(&long);
assert!(result.len() <= 2006); // 1990 content + 16 suffix
assert!(result.ends_with("\n... (truncated)"));
}
#[test]
fn test_broadcast_dm_validates_snowflake() {
// broadcast_dm rejects invalid Discord snowflake IDs before making
// any API calls. We can call it directly since invalid IDs are
// rejected before any host function is invoked.
assert!(broadcast_dm("", "hi").is_err());
assert!(broadcast_dm("abc", "hi").is_err());
assert!(broadcast_dm("12345", "hi").is_err()); // too short
assert!(broadcast_dm("123456789012345678901", "hi").is_err()); // too long
assert!(broadcast_dm("12345678901234567x", "hi").is_err()); // non-digit
}
}
@@ -1,4 +0,0 @@
-- Add source_channel to conversations for cross-channel approval authorization.
-- Tracks which channel originally created a conversation so that approval
-- messages from other channels can be validated.
ALTER TABLE conversations ADD COLUMN source_channel TEXT;
+2 -35
View File
@@ -843,10 +843,7 @@ impl Agent {
{
use crate::agent::session::Thread;
let mut sess = session.lock().await;
// Bootstrap thread has no incoming message -- use the
// "__bootstrap__" sentinel so approvals from any channel are
// permitted. None means "deny by default" (fail-closed).
let thread = Thread::with_id(id, sess.id, Some("__bootstrap__"));
let thread = Thread::with_id(id, sess.id);
sess.active_thread = Some(id);
sess.threads.entry(id).or_insert(thread);
}
@@ -1151,37 +1148,7 @@ impl Agent {
.get_or_create_session(&message.user_id)
.await;
let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get(&target_thread_id) {
// Verify the thread actually has a pending approval before
// allowing approval-shaped messages to target it. Without this
// check, an attacker could use approval messages to hijack any
// thread by UUID.
if thread.pending_approval.is_none() {
tracing::warn!(
%target_thread_id,
approval_channel = %message.channel,
"Blocked approval for thread with no pending approval"
);
drop(sess);
return Ok(Some("Error: no pending approval on this thread".into()));
}
let authorized = crate::agent::session::is_approval_authorized(
thread.source_channel.as_deref(),
&message.channel,
);
if !authorized {
tracing::warn!(
%target_thread_id,
source_channel = ?thread.source_channel,
approval_channel = %message.channel,
"Blocked cross-channel approval attempt"
);
drop(sess);
return Ok(Some(
"Error: approval not authorized for this channel".into(),
));
}
if sess.threads.contains_key(&target_thread_id) {
sess.active_thread = Some(target_thread_id);
sess.last_active_at = chrono::Utc::now();
drop(sess);
+5 -5
View File
@@ -319,7 +319,7 @@ mod tests {
#[test]
fn test_format_turns() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("Hello");
thread.complete_turn("Hi there");
thread.start_turn("How are you?");
@@ -351,7 +351,7 @@ mod tests {
/// Helper: build a thread with `n` completed turns.
/// Turn `i` has user_input "msg-{i}" and response "resp-{i}".
fn make_thread(n: usize) -> Thread {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
for i in 0..n {
thread.start_turn(format!("msg-{}", i));
thread.complete_turn(format!("resp-{}", i));
@@ -457,7 +457,7 @@ mod tests {
async fn test_compact_truncate_empty_turns() {
let llm = Arc::new(StubLlm::new("unused"));
let compactor = make_compactor(llm);
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
assert!(thread.turns.is_empty());
let result = compactor
@@ -698,7 +698,7 @@ mod tests {
#[test]
fn test_format_turns_for_storage_with_tool_calls() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("Search for X");
// Record a tool call on the current turn
if let Some(turn) = thread.turns.last_mut() {
@@ -719,7 +719,7 @@ mod tests {
#[test]
fn test_format_turns_for_storage_incomplete_turn() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("In progress message");
// Don't complete the turn
+2 -2
View File
@@ -2299,7 +2299,7 @@ mod tests {
// Initialize a thread in the session so the loop can record tool calls.
let thread_id = {
let mut sess = session.lock().await;
sess.create_thread(Some("test")).id
sess.create_thread().id
};
let message = IncomingMessage::new("test", "test-user", "do something");
@@ -2412,7 +2412,7 @@ mod tests {
let session = Arc::new(Mutex::new(Session::new("test-user")));
let thread_id = {
let mut sess = session.lock().await;
sess.create_thread(Some("test")).id
sess.create_thread().id
};
let message = IncomingMessage::new("test", "test-user", "keep calling tools");
+244 -46
View File
@@ -1089,48 +1089,134 @@ async fn execute_routine(ctx: EngineContext, routine: Routine, run: RoutineRun)
// Increment running count (atomic: survives panics in the execution below)
ctx.running_count.fetch_add(1, Ordering::Relaxed);
let result = match &routine.action {
RoutineAction::Lightweight {
prompt,
context_paths,
max_tokens,
use_tools,
max_tool_rounds,
} => {
execute_lightweight(
&ctx,
&routine,
prompt,
context_paths,
*max_tokens,
*use_tools,
*max_tool_rounds,
)
.await
// Retry constants for transient lightweight execution failures.
const MAX_RETRIES: u32 = 3;
const BASE_DELAY_MS: u64 = 1000;
let is_lightweight = matches!(routine.action, RoutineAction::Lightweight { .. });
// The retry block returns both the execution result and any accumulated
// token count so that usage is preserved even on final failure.
let (result, accumulated_tokens) = {
let mut attempt = 0u32;
// Track accumulated tokens as Option to preserve None semantics:
// None = no attempt reported tokens; Some(n) = at least one attempt did.
let mut accumulated_tokens: Option<i32> = None;
let uses_tools = matches!(
routine.action,
RoutineAction::Lightweight {
use_tools: true,
..
}
) && ctx.config.lightweight_tools_enabled;
/// Extract partial_tokens from any RoutineError variant that carries them.
fn extract_partial_tokens(e: &RoutineError) -> Option<i32> {
match e {
RoutineError::LlmFailed {
partial_tokens: Some(t),
..
}
| RoutineError::EmptyResponse {
partial_tokens: Some(t),
}
| RoutineError::TruncatedResponse {
partial_tokens: Some(t),
} => Some(*t),
_ => None,
}
}
RoutineAction::FullJob {
title,
description,
max_iterations,
} => {
let execution = FullJobExecutionConfig {
title,
description,
max_iterations: *max_iterations,
/// Merge an optional partial token count into the accumulator,
/// only materializing Some when at least one source had Some.
fn accumulate(acc: Option<i32>, partial: Option<i32>) -> Option<i32> {
match (acc, partial) {
(Some(a), Some(p)) => Some(a.saturating_add(p)),
(Some(a), None) => Some(a),
(None, p) => p,
}
}
loop {
let execution_result = match &routine.action {
RoutineAction::Lightweight {
prompt,
context_paths,
max_tokens,
use_tools,
max_tool_rounds,
} => {
execute_lightweight(
&ctx,
&routine,
prompt,
context_paths,
*max_tokens,
*use_tools,
*max_tool_rounds,
)
.await
}
RoutineAction::FullJob {
title,
description,
max_iterations,
} => {
let execution = FullJobExecutionConfig {
title,
description,
max_iterations: *max_iterations,
};
execute_full_job(&ctx, &routine, &run, &execution).await
}
};
execute_full_job(&ctx, &routine, &run, &execution).await
match execution_result {
Ok((status, summary, tokens)) => {
// Merge tokens: only produce Some when at least one source had Some.
let total = accumulate(accumulated_tokens, tokens);
break (Ok((status, summary, total)), accumulated_tokens);
}
Err(ref e)
if is_lightweight
&& !uses_tools
&& e.is_retryable()
// Skip outer retry for LlmFailed — RetryProvider already
// retries transient LLM errors with its own budget. Retrying
// here would create a multiplicative retry count.
&& !matches!(e, RoutineError::LlmFailed { .. })
&& attempt < MAX_RETRIES =>
{
// Accumulate partial tokens from the failed attempt.
accumulated_tokens = accumulate(accumulated_tokens, extract_partial_tokens(e));
attempt += 1;
let delay = Duration::from_millis(
BASE_DELAY_MS.saturating_mul(2u64.saturating_pow(attempt - 1)),
);
tracing::event!(target: "transient_routine_errors", tracing::Level::WARN, routine = %routine.name, attempt = attempt, max_retries = MAX_RETRIES, delay_ms = delay.as_millis() as u64, "Transient routine error, retrying: {}", e);
tokio::time::sleep(delay).await;
}
Err(e) => {
// Accumulate tokens from the final failed attempt.
accumulated_tokens = accumulate(accumulated_tokens, extract_partial_tokens(&e));
break (Err(e), accumulated_tokens);
}
}
}
};
// Decrement running count
ctx.running_count.fetch_sub(1, Ordering::Relaxed);
// Process result
// Process result — on failure, preserve accumulated token total from
// earlier retry attempts so usage reporting stays accurate.
let (status, summary, tokens) = match result {
Ok(execution) => execution,
Err(e) => {
tracing::error!(routine = %routine.name, "Execution failed: {}", e);
(RunStatus::Failed, Some(e.to_string()), None)
(RunStatus::Failed, Some(e.to_string()), accumulated_tokens)
}
};
@@ -1511,13 +1597,14 @@ async fn execute_lightweight_no_tools(
.with_max_tokens(effective_max_tokens)
.with_temperature(0.3);
let response = ctx
.llm
.complete(request)
.await
.map_err(|e| RoutineError::LlmFailed {
let response = ctx.llm.complete(request).await.map_err(|e| {
let retryable = crate::llm::retry::is_retryable(&e);
RoutineError::LlmFailed {
reason: e.to_string(),
})?;
partial_tokens: None,
retryable,
}
})?;
handle_text_response(
&response.content,
@@ -1538,12 +1625,18 @@ fn handle_text_response(
) -> Result<(RunStatus, Option<String>, Option<i32>), RoutineError> {
let content = content.trim();
// Empty content guard
// Empty content guard — carry consumed tokens so the retry loop can
// accumulate them even when the response shape is invalid.
if content.is_empty() {
let consumed = Some((total_input_tokens + total_output_tokens) as i32);
return if finish_reason == FinishReason::Length {
Err(RoutineError::TruncatedResponse)
Err(RoutineError::TruncatedResponse {
partial_tokens: consumed,
})
} else {
Err(RoutineError::EmptyResponse)
Err(RoutineError::EmptyResponse {
partial_tokens: consumed,
})
};
}
@@ -1621,13 +1714,15 @@ async fn execute_lightweight_with_tools(
.with_max_tokens(effective_max_tokens)
.with_temperature(0.3);
let response =
ctx.llm
.complete(request)
.await
.map_err(|e| RoutineError::LlmFailed {
reason: e.to_string(),
})?;
let response = ctx.llm.complete(request).await.map_err(|e| {
let partial = (total_input_tokens + total_output_tokens) as i32;
let retryable = crate::llm::retry::is_retryable(&e);
RoutineError::LlmFailed {
reason: e.to_string(),
partial_tokens: if partial > 0 { Some(partial) } else { None },
retryable,
}
})?;
total_input_tokens += response.input_tokens;
total_output_tokens += response.output_tokens;
@@ -1654,8 +1749,12 @@ async fn execute_lightweight_with_tools(
.with_temperature(0.3);
let response = ctx.llm.complete_with_tools(request).await.map_err(|e| {
let partial = (total_input_tokens + total_output_tokens) as i32;
let retryable = crate::llm::retry::is_retryable(&e);
RoutineError::LlmFailed {
reason: e.to_string(),
partial_tokens: if partial > 0 { Some(partial) } else { None },
retryable,
}
})?;
@@ -1828,6 +1927,7 @@ async fn execute_routine_tool(
}
/// Send a notification based on the routine's notify config and run status.
#[allow(clippy::too_many_arguments)]
async fn send_notification(
tx: &mpsc::Sender<OutgoingResponse>,
notify: &NotifyConfig,
@@ -2513,6 +2613,104 @@ mod tests {
}
}
/// Regression test for #1320: transient errors are retried for lightweight
/// routines but not for full-job routines or hard failures.
#[test]
fn test_retry_classification_for_routine_errors() {
use crate::error::RoutineError;
// Transient errors (retryable for lightweight routines)
let transient_errors: Vec<RoutineError> = vec![
RoutineError::LlmFailed {
reason: "rate limit".into(),
partial_tokens: None,
retryable: true,
},
RoutineError::LlmFailed {
reason: "network timeout".into(),
partial_tokens: Some(42),
retryable: true,
},
RoutineError::EmptyResponse {
partial_tokens: None,
},
RoutineError::TruncatedResponse {
partial_tokens: Some(100),
},
];
for err in &transient_errors {
assert!(err.is_retryable(), "{} should be retryable", err);
}
// Permanent LLM failures that should NOT be retried
// (retryable: false is set at conversion time by llm::retry::is_retryable)
let permanent_llm_errors: Vec<RoutineError> = vec![
RoutineError::LlmFailed {
reason: "Authentication failed for provider openai".into(),
partial_tokens: None,
retryable: false,
},
RoutineError::LlmFailed {
reason: "invalid_api_key: bad key".into(),
partial_tokens: None,
retryable: false,
},
RoutineError::LlmFailed {
reason: "content policy violation".into(),
partial_tokens: None,
retryable: false,
},
RoutineError::LlmFailed {
reason: "content_filter triggered".into(),
partial_tokens: None,
retryable: false,
},
RoutineError::LlmFailed {
reason: "context length exceeded: 150000 tokens used, 128000 allowed".into(),
partial_tokens: Some(100),
retryable: false,
},
RoutineError::LlmFailed {
reason: "model not available on provider anthropic".into(),
partial_tokens: None,
retryable: false,
},
RoutineError::LlmFailed {
reason: "content moderation flagged".into(),
partial_tokens: None,
retryable: false,
},
];
for err in &permanent_llm_errors {
assert!(!err.is_retryable(), "{} should NOT be retryable", err);
}
// Hard failures (never retried)
let hard_errors: Vec<RoutineError> = vec![
RoutineError::Disabled {
name: "test".into(),
},
RoutineError::NotFound {
id: uuid::Uuid::new_v4(),
},
RoutineError::NotAuthorized {
id: uuid::Uuid::new_v4(),
},
RoutineError::MaxConcurrent {
name: "test".into(),
},
RoutineError::JobDispatchFailed {
reason: "no docker".into(),
},
RoutineError::Database {
reason: "connection refused".into(),
},
];
for err in &hard_errors {
assert!(!err.is_retryable(), "{} should NOT be retryable", err);
}
}
#[test]
fn test_sanitize_summary_strips_control_chars() {
use super::sanitize_summary;
+54 -169
View File
@@ -68,8 +68,8 @@ impl Session {
}
/// Create a new thread in this session.
pub fn create_thread(&mut self, channel: Option<&str>) -> &mut Thread {
let thread = Thread::new(self.id, channel);
pub fn create_thread(&mut self) -> &mut Thread {
let thread = Thread::new(self.id);
let thread_id = thread.id;
self.active_thread = Some(thread_id);
self.last_active_at = Utc::now();
@@ -87,9 +87,9 @@ impl Session {
}
/// Get or create the active thread.
pub fn get_or_create_thread(&mut self, channel: Option<&str>) -> &mut Thread {
pub fn get_or_create_thread(&mut self) -> &mut Thread {
match self.active_thread {
None => self.create_thread(channel),
None => self.create_thread(),
Some(id) => {
if self.threads.contains_key(&id) {
// Entry existence confirmed by contains_key above.
@@ -100,7 +100,7 @@ impl Session {
} else {
// Stale active_thread ID: create a new thread, which
// updates self.active_thread to the new thread's ID.
self.create_thread(channel)
self.create_thread()
}
}
}
@@ -225,9 +225,6 @@ pub struct Thread {
/// Messages queued while the thread was processing a turn.
#[serde(default, skip_serializing_if = "VecDeque::is_empty")]
pub pending_messages: VecDeque<String>,
/// Channel that created this thread (for approval authorization).
#[serde(default)]
pub source_channel: Option<String>,
}
/// Maximum number of messages that can be queued while a thread is processing.
@@ -236,29 +233,9 @@ pub struct Thread {
/// rapid follow-ups. The drain loop processes them as one newline-delimited turn.
pub const MAX_PENDING_MESSAGES: usize = 10;
/// Sentinel value for bootstrap threads that accept approvals from any channel.
pub const BOOTSTRAP_SOURCE_CHANNEL: &str = "__bootstrap__";
/// Check whether an approval from `requesting_channel` is authorized for a
/// thread whose `source_channel` is `source`.
///
/// Rules:
/// - `None` (unknown origin) -> denied (fail-closed)
/// - `Some("__bootstrap__")` -> authorized from any channel
/// - `Some(src) == requesting` -> same channel, authorized
/// - requesting is "web" or "gateway" -> always authorized (trusted UI)
/// - Otherwise -> denied
pub fn is_approval_authorized(source: Option<&str>, requesting: &str) -> bool {
match source {
None => false,
Some(src) if src == BOOTSTRAP_SOURCE_CHANNEL => true,
Some(src) => src == requesting || requesting == "web" || requesting == "gateway",
}
}
impl Thread {
/// Create a new thread.
pub fn new(session_id: Uuid, source_channel: Option<&str>) -> Self {
pub fn new(session_id: Uuid) -> Self {
let now = Utc::now();
Self {
id: Uuid::new_v4(),
@@ -271,12 +248,11 @@ impl Thread {
pending_approval: None,
pending_auth: None,
pending_messages: VecDeque::new(),
source_channel: source_channel.map(String::from),
}
}
/// Create a thread with a specific ID (for DB hydration).
pub fn with_id(id: Uuid, session_id: Uuid, source_channel: Option<&str>) -> Self {
pub fn with_id(id: Uuid, session_id: Uuid) -> Self {
let now = Utc::now();
Self {
id,
@@ -289,7 +265,6 @@ impl Thread {
pending_approval: None,
pending_auth: None,
pending_messages: VecDeque::new(),
source_channel: source_channel.map(String::from),
}
}
@@ -812,13 +787,13 @@ mod tests {
let mut session = Session::new("user-123");
assert!(session.active_thread.is_none());
session.create_thread(None);
session.create_thread();
assert!(session.active_thread.is_some());
}
#[test]
fn test_thread_turns() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("Hello");
assert_eq!(thread.state, ThreadState::Processing);
@@ -831,7 +806,7 @@ mod tests {
#[test]
fn test_thread_messages() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("First message");
thread.complete_turn("First response");
@@ -854,7 +829,7 @@ mod tests {
#[test]
fn test_restore_from_messages() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// First add some turns
thread.start_turn("Original message");
@@ -880,7 +855,7 @@ mod tests {
#[test]
fn test_restore_from_messages_incomplete_turn() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Messages with incomplete last turn (no assistant response)
let messages = vec![
@@ -899,7 +874,7 @@ mod tests {
#[test]
fn test_enter_auth_mode() {
let before = Utc::now();
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
assert!(thread.pending_auth.is_none());
thread.enter_auth_mode("telegram".to_string());
@@ -912,7 +887,7 @@ mod tests {
#[test]
fn test_take_pending_auth() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.enter_auth_mode("notion".to_string());
let pending = thread.take_pending_auth();
@@ -927,7 +902,7 @@ mod tests {
#[test]
fn test_pending_auth_serialization() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.enter_auth_mode("openai".to_string());
let json = serde_json::to_string(&thread).expect("should serialize");
@@ -957,7 +932,7 @@ mod tests {
#[test]
fn test_pending_auth_default_none() {
// Deserialization of old data without pending_auth should default to None
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.pending_auth = None;
let json = serde_json::to_string(&thread).expect("serialize");
@@ -971,7 +946,7 @@ mod tests {
fn test_thread_with_id() {
let specific_id = Uuid::new_v4();
let session_id = Uuid::new_v4();
let thread = Thread::with_id(specific_id, session_id, None);
let thread = Thread::with_id(specific_id, session_id);
assert_eq!(thread.id, specific_id);
assert_eq!(thread.session_id, session_id);
@@ -983,7 +958,7 @@ mod tests {
fn test_thread_with_id_restore_messages() {
let thread_id = Uuid::new_v4();
let session_id = Uuid::new_v4();
let mut thread = Thread::with_id(thread_id, session_id, None);
let mut thread = Thread::with_id(thread_id, session_id);
let messages = vec![
ChatMessage::user("Hello from DB"),
@@ -1002,7 +977,7 @@ mod tests {
#[test]
fn test_restore_from_messages_empty() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Add a turn first, then restore with empty vec
thread.start_turn("hello");
@@ -1018,7 +993,7 @@ mod tests {
#[test]
fn test_restore_from_messages_only_assistant_messages() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Only assistant messages (no user messages to anchor turns)
let messages = vec![
@@ -1035,7 +1010,7 @@ mod tests {
#[test]
fn test_restore_from_messages_multiple_user_messages_in_a_row() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Two user messages with no assistant response between them
let messages = vec![
@@ -1062,8 +1037,8 @@ mod tests {
fn test_thread_switch() {
let mut session = Session::new("user-1");
let t1_id = session.create_thread(None).id;
let t2_id = session.create_thread(None).id;
let t1_id = session.create_thread().id;
let t2_id = session.create_thread().id;
// After creating two threads, active should be the last one
assert_eq!(session.active_thread, Some(t2_id));
@@ -1083,8 +1058,8 @@ mod tests {
fn test_get_or_create_thread_idempotent() {
let mut session = Session::new("user-1");
let tid1 = session.get_or_create_thread(None).id;
let tid2 = session.get_or_create_thread(None).id;
let tid1 = session.get_or_create_thread().id;
let tid2 = session.get_or_create_thread().id;
// Should return the same thread (not create a new one each time)
assert_eq!(tid1, tid2);
@@ -1093,7 +1068,7 @@ mod tests {
#[test]
fn test_truncate_turns() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
for i in 0..5 {
thread.start_turn(format!("msg-{}", i));
@@ -1117,7 +1092,7 @@ mod tests {
#[test]
fn test_truncate_turns_noop_when_fewer() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("only one");
thread.complete_turn("response");
@@ -1129,7 +1104,7 @@ mod tests {
#[test]
fn test_thread_interrupt_and_resume() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("do something");
assert_eq!(thread.state, ThreadState::Processing);
@@ -1147,7 +1122,7 @@ mod tests {
#[test]
fn test_resume_only_from_interrupted() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Idle thread: resume should be a no-op
assert_eq!(thread.state, ThreadState::Idle);
@@ -1163,7 +1138,7 @@ mod tests {
#[test]
fn test_turn_fail() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("risky operation");
thread.fail_turn("connection timed out");
@@ -1179,7 +1154,7 @@ mod tests {
#[test]
fn test_messages_with_incomplete_last_turn() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("first");
thread.complete_turn("first reply");
@@ -1195,7 +1170,7 @@ mod tests {
#[test]
fn test_thread_serialization_round_trip() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("hello");
thread.complete_turn("world");
@@ -1213,7 +1188,7 @@ mod tests {
#[test]
fn test_session_serialization_round_trip() {
let mut session = Session::new("user-ser");
session.create_thread(None);
session.create_thread();
session.auto_approve_tool("echo");
let json = serde_json::to_string(&session).unwrap();
@@ -1251,7 +1226,7 @@ mod tests {
#[test]
fn test_turn_number_increments() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Before any turns, turn_number() is 1 (1-indexed for display)
assert_eq!(thread.turn_number(), 1);
@@ -1266,7 +1241,7 @@ mod tests {
#[test]
fn test_complete_turn_on_empty_thread() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Completing a turn when there are no turns should be a safe no-op
thread.complete_turn("phantom response");
@@ -1276,7 +1251,7 @@ mod tests {
#[test]
fn test_fail_turn_on_empty_thread() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Failing a turn when there are no turns should be a safe no-op
thread.fail_turn("phantom error");
@@ -1286,7 +1261,7 @@ mod tests {
#[test]
fn test_pending_approval_flow() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
let approval = PendingApproval {
request_id: Uuid::new_v4(),
@@ -1313,7 +1288,7 @@ mod tests {
#[test]
fn test_clear_pending_approval() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
let approval = PendingApproval {
request_id: Uuid::new_v4(),
@@ -1342,7 +1317,7 @@ mod tests {
assert!(session.active_thread().is_none());
assert!(session.active_thread_mut().is_none());
let tid = session.create_thread(None).id;
let tid = session.create_thread().id;
assert!(session.active_thread().is_some());
assert_eq!(session.active_thread().unwrap().id, tid);
@@ -1359,7 +1334,7 @@ mod tests {
#[test]
fn test_messages_includes_tool_calls() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("Search for X");
{
@@ -1391,7 +1366,7 @@ mod tests {
#[test]
fn test_messages_multiple_tool_calls_per_turn() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("Do two things");
{
@@ -1418,7 +1393,7 @@ mod tests {
#[test]
fn test_restore_from_messages_with_tool_calls() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Build a message sequence with tool calls
let tc = ToolCall {
@@ -1450,7 +1425,7 @@ mod tests {
#[test]
fn test_restore_from_messages_with_tool_error() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
let tc = ToolCall {
id: "call_0".to_string(),
@@ -1481,7 +1456,7 @@ mod tests {
fn test_messages_round_trip_with_tools() {
// Build a thread with tool calls, get messages(), restore, get messages() again
// The two message sequences should be equivalent.
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("Do search");
{
@@ -1494,7 +1469,7 @@ mod tests {
let messages_original = thread.messages();
// Restore into a new thread
let mut thread2 = Thread::new(Uuid::new_v4(), None);
let mut thread2 = Thread::new(Uuid::new_v4());
thread2.restore_from_messages(messages_original.clone());
let messages_restored = thread2.messages();
@@ -1516,7 +1491,7 @@ mod tests {
#[test]
fn test_restore_multi_stage_tool_calls() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
let tc1 = ToolCall {
id: "call_a".to_string(),
@@ -1559,7 +1534,7 @@ mod tests {
#[test]
fn test_messages_truncates_large_tool_results() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("Read big file");
{
@@ -1582,7 +1557,7 @@ mod tests {
#[test]
fn test_thread_message_queue() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Queue is initially empty
assert!(thread.pending_messages.is_empty());
@@ -1618,7 +1593,7 @@ mod tests {
#[test]
fn test_thread_message_queue_serialization() {
let mut thread = Thread::new(Uuid::new_v4(), None);
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();
@@ -1638,7 +1613,7 @@ mod tests {
#[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(), None);
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
@@ -1649,7 +1624,7 @@ mod tests {
#[test]
fn test_interrupt_clears_pending_messages() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Start a turn so there's something to interrupt
thread.start_turn("initial input");
@@ -1668,7 +1643,7 @@ mod tests {
#[test]
fn test_thread_state_idle_after_full_drain() {
let mut thread = Thread::new(Uuid::new_v4(), None);
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).
@@ -1696,7 +1671,7 @@ mod tests {
#[test]
fn test_drain_pending_messages_merges_with_newlines() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Empty queue returns None
assert!(thread.drain_pending_messages().is_none());
@@ -1725,7 +1700,7 @@ mod tests {
#[test]
fn test_requeue_drained_preserves_content_at_front() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Re-queue into empty queue
thread.requeue_drained("failed batch".to_string());
@@ -1836,94 +1811,4 @@ mod tests {
&serde_json::json!("done")
);
}
#[test]
fn test_thread_new_stores_source_channel() {
let thread = Thread::new(Uuid::new_v4(), Some("telegram"));
assert_eq!(thread.source_channel.as_deref(), Some("telegram"));
}
#[test]
fn test_thread_new_none_channel() {
let thread = Thread::new(Uuid::new_v4(), None);
assert!(thread.source_channel.is_none());
}
#[test]
fn test_source_channel_serde_backcompat() {
// Simulate deserializing a Thread from older DB records that lack source_channel.
let thread = Thread::new(Uuid::new_v4(), Some("cli"));
let json = serde_json::to_string(&thread).unwrap();
// Remove the source_channel field to simulate an old record.
let mut value: serde_json::Value = serde_json::from_str(&json).unwrap();
value.as_object_mut().unwrap().remove("source_channel");
let old_json = serde_json::to_string(&value).unwrap();
let deserialized: Thread = serde_json::from_str(&old_json).unwrap();
assert!(
deserialized.source_channel.is_none(),
"missing source_channel should deserialize as None"
);
}
#[test]
fn test_approval_authorized_same_channel() {
assert!(
is_approval_authorized(Some("telegram"), "telegram"),
"same channel should be authorized"
);
}
#[test]
fn test_approval_authorized_different_channel_blocked() {
assert!(
!is_approval_authorized(Some("telegram"), "http"),
"different channel should be blocked"
);
}
#[test]
fn test_approval_authorized_web_always_allowed() {
assert!(
is_approval_authorized(Some("telegram"), "web"),
"web channel should always be authorized"
);
}
#[test]
fn test_approval_authorized_gateway_always_allowed() {
assert!(
is_approval_authorized(Some("telegram"), "gateway"),
"gateway channel should always be authorized"
);
}
#[test]
fn test_approval_authorized_none_denied() {
assert!(
!is_approval_authorized(None, "telegram"),
"None source_channel should be denied (fail-closed)"
);
assert!(
!is_approval_authorized(None, "web"),
"None source_channel should be denied even for web"
);
}
#[test]
fn test_approval_authorized_bootstrap_any_channel() {
assert!(
is_approval_authorized(Some(BOOTSTRAP_SOURCE_CHANNEL), "telegram"),
"__bootstrap__ should be authorized from any channel"
);
assert!(
is_approval_authorized(Some(BOOTSTRAP_SOURCE_CHANNEL), "http"),
"__bootstrap__ should be authorized from any channel"
);
assert!(
is_approval_authorized(Some(BOOTSTRAP_SOURCE_CHANNEL), "cli"),
"__bootstrap__ should be authorized from any channel"
);
}
}
+13 -28
View File
@@ -200,7 +200,7 @@ impl SessionManager {
// Create new thread (always create a new one for a new key)
let thread_id = {
let mut sess = session.lock().await;
let thread = sess.create_thread(Some(channel));
let thread = sess.create_thread();
thread.id
};
@@ -476,7 +476,7 @@ mod tests {
let session = Arc::new(Mutex::new(Session::new("user-hydrate")));
{
let mut sess = session.lock().await;
let thread = Thread::with_id(thread_id, sess.id, None);
let thread = Thread::with_id(thread_id, sess.id);
sess.threads.insert(thread_id, thread);
sess.active_thread = Some(thread_id);
}
@@ -600,7 +600,7 @@ mod tests {
// Simulate hydration: create thread with a known UUID
{
let mut sess = session.lock().await;
let thread = Thread::with_id(known_uuid, session_id, None);
let thread = Thread::with_id(known_uuid, session_id);
sess.threads.insert(known_uuid, thread);
}
@@ -627,7 +627,7 @@ mod tests {
let session = Arc::new(Mutex::new(Session::new("user-idem")));
{
let mut sess = session.lock().await;
let thread = Thread::with_id(tid, sess.id, None);
let thread = Thread::with_id(tid, sess.id);
sess.threads.insert(tid, thread);
}
@@ -656,7 +656,7 @@ mod tests {
let session = Arc::new(Mutex::new(Session::new("user-undo")));
{
let mut sess = session.lock().await;
let thread = Thread::with_id(tid, sess.id, None);
let thread = Thread::with_id(tid, sess.id);
sess.threads.insert(tid, thread);
}
@@ -680,7 +680,7 @@ mod tests {
let session = Arc::new(Mutex::new(Session::new("user-new")));
{
let mut sess = session.lock().await;
let thread = Thread::with_id(tid, sess.id, None);
let thread = Thread::with_id(tid, sess.id);
sess.threads.insert(tid, thread);
}
@@ -788,7 +788,7 @@ mod tests {
let session = Arc::new(Mutex::new(Session::new("user-cross")));
{
let mut sess = session.lock().await;
let thread = Thread::with_id(tid, sess.id, None);
let thread = Thread::with_id(tid, sess.id);
sess.threads.insert(tid, thread);
}
@@ -815,7 +815,7 @@ mod tests {
let session = Arc::new(Mutex::new(Session::new("user-cross")));
{
let mut sess = session.lock().await;
let thread = Thread::with_id(tid, sess.id, None);
let thread = Thread::with_id(tid, sess.id);
sess.threads.insert(tid, thread);
}
@@ -966,7 +966,7 @@ mod tests {
let adopted_id = Uuid::new_v4();
{
let mut sess = session1.lock().await;
let thread = Thread::with_id(adopted_id, sess.id, None);
let thread = Thread::with_id(adopted_id, sess.id);
sess.threads.insert(adopted_id, thread);
}
// Resolve with the UUID as external_thread_id -- should adopt it
@@ -992,7 +992,7 @@ mod tests {
let session = Arc::new(Mutex::new(Session::new("user-direct")));
{
let mut sess = session.lock().await;
let thread = Thread::with_id(tid, sess.id, None);
let thread = Thread::with_id(tid, sess.id);
sess.threads.insert(tid, thread);
}
{
@@ -1030,7 +1030,7 @@ mod tests {
let known_id = Uuid::new_v4();
{
let mut sess = session.lock().await;
let thread = Thread::with_id(known_id, sess.id, None);
let thread = Thread::with_id(known_id, sess.id);
sess.threads.insert(known_id, thread);
}
@@ -1057,7 +1057,7 @@ mod tests {
let known_id = Uuid::new_v4();
{
let mut sess = session.lock().await;
let thread = Thread::with_id(known_id, sess.id, None);
let thread = Thread::with_id(known_id, sess.id);
sess.threads.insert(known_id, thread);
}
@@ -1081,7 +1081,7 @@ mod tests {
let known_id = Uuid::new_v4();
{
let mut sess = session.lock().await;
let thread = Thread::with_id(known_id, sess.id, None);
let thread = Thread::with_id(known_id, sess.id);
sess.threads.insert(known_id, thread);
}
@@ -1102,19 +1102,4 @@ mod tests {
"should NOT adopt UUID when external_thread_id is None"
);
}
#[tokio::test]
async fn test_thread_stores_source_channel() {
let manager = SessionManager::new();
let (session, thread_id) = manager.resolve_thread("user-1", "telegram", None).await;
let sess = session.lock().await;
let thread = sess.threads.get(&thread_id).unwrap();
assert_eq!(
thread.source_channel.as_deref(),
Some("telegram"),
"resolve_thread should store source_channel from the channel parameter"
);
}
}
+9 -25
View File
@@ -135,29 +135,13 @@ impl Agent {
msg_count = 0;
}
// Create thread with the historical ID and restore messages.
// Read source_channel from DB so the authorization check uses the
// original creator's channel, not the requesting message's channel.
let db_source_channel = if let Some(store) = self.store() {
store
.get_conversation_source_channel(thread_uuid)
.await
.unwrap_or(None)
} else {
None
};
let effective_source_channel = db_source_channel.as_deref();
// Create thread with the historical ID and restore messages
let session_id = {
let sess = session.lock().await;
sess.id
};
let mut thread = crate::agent::session::Thread::with_id(
thread_uuid,
session_id,
effective_source_channel,
);
let mut thread = crate::agent::session::Thread::with_id(thread_uuid, session_id);
if !chat_messages.is_empty() {
thread.restore_from_messages(chat_messages);
}
@@ -652,7 +636,7 @@ impl Agent {
user_id: &str,
) -> bool {
match store
.ensure_conversation(thread_id, channel, user_id, None, Some(channel))
.ensure_conversation(thread_id, channel, user_id, None)
.await
{
Ok(true) => true,
@@ -1797,7 +1781,7 @@ impl Agent {
.get_or_create_session(&message.user_id)
.await;
let mut sess = session.lock().await;
let thread = sess.create_thread(Some(&message.channel));
let thread = sess.create_thread();
let thread_id = thread.id;
Ok(SubmissionResult::ok_with_message(format!(
"New thread: {}",
@@ -2133,7 +2117,7 @@ mod tests {
let session_id = Uuid::new_v4();
let thread_id = Uuid::new_v4();
let mut thread = Thread::with_id(thread_id, session_id, None);
let mut thread = Thread::with_id(thread_id, session_id);
// Set thread to AwaitingApproval with a pending tool approval
let pending = PendingApproval {
@@ -2201,7 +2185,7 @@ mod tests {
use crate::agent::session::{MAX_PENDING_MESSAGES, Thread, ThreadState};
use uuid::Uuid;
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("processing something");
assert_eq!(thread.state, ThreadState::Processing);
@@ -2227,7 +2211,7 @@ mod tests {
use crate::agent::session::{Thread, ThreadState};
use uuid::Uuid;
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("processing");
thread.queue_message("pending-1".to_string());
@@ -2257,7 +2241,7 @@ mod tests {
let thread_id = Uuid::new_v4();
let session_id = Uuid::new_v4();
let mut thread = Thread::with_id(thread_id, session_id, None);
let mut thread = Thread::with_id(thread_id, session_id);
thread.start_turn("working");
assert_eq!(thread.state, ThreadState::Processing);
@@ -2284,7 +2268,7 @@ mod tests {
let thread_id = Uuid::new_v4();
let session_id = Uuid::new_v4();
let mut thread = Thread::with_id(thread_id, session_id, None);
let mut thread = Thread::with_id(thread_id, session_id);
thread.start_turn("working");
assert_eq!(thread.state, ThreadState::Processing);
-40
View File
@@ -34,7 +34,6 @@ pub async fn setup_wasm_channels(
secrets_store: &Option<Arc<dyn SecretsStore + Send + Sync>>,
extension_manager: Option<&Arc<ExtensionManager>>,
database: Option<&Arc<dyn Database>>,
registered_channel_names: &[String],
) -> Option<WasmChannelSetup> {
let runtime = match WasmChannelRuntime::new(WasmChannelRuntimeConfig::default()) {
Ok(r) => Arc::new(r),
@@ -72,46 +71,7 @@ pub async fn setup_wasm_channels(
let mut channels: Vec<(String, Box<dyn crate::channels::Channel>)> = Vec::new();
let mut channel_names: Vec<String> = Vec::new();
// Reserved channel names that WASM modules must not claim.
// A malicious module could otherwise register as a trusted built-in
// channel and bypass cross-channel authorization checks.
// This list must cover every built-in channel name to prevent a WASM
// module from impersonating a built-in and satisfying same-channel
// approval checks.
const RESERVED_CHANNEL_NAMES: &[&str] = &[
"web",
"gateway",
"cli",
"repl",
"http",
"signal",
"slack-relay",
"secret_save",
];
for loaded in results.loaded {
let name_lower = loaded.name().to_ascii_lowercase();
if RESERVED_CHANNEL_NAMES.contains(&name_lower.as_str()) {
tracing::warn!(
channel = %loaded.name(),
"Rejected WASM channel with reserved name"
);
continue;
}
// Also reject any name that collides with an already-registered
// channel to prevent a WASM module from shadowing a channel that
// was registered earlier in the startup sequence.
if registered_channel_names
.iter()
.any(|n| n.to_ascii_lowercase() == name_lower)
{
tracing::warn!(
channel = %loaded.name(),
"Rejected WASM channel that collides with already-registered channel"
);
continue;
}
let (name, channel) = register_channel(
loaded,
config,
+2 -8
View File
@@ -574,7 +574,7 @@ pub async fn chat_new_thread_handler(
.await;
let (thread_id, info) = {
let mut sess = session.lock().await;
let thread = sess.create_thread(Some("web"));
let thread = sess.create_thread();
let id = thread.id;
let info = ThreadInfo {
id: thread.id,
@@ -593,13 +593,7 @@ pub async fn chat_new_thread_handler(
// so that the subsequent loadThreads() call from the frontend sees it.
if let Some(ref store) = state.store {
match store
.ensure_conversation(
thread_id,
"gateway",
&identity.user_id,
None,
Some("gateway"),
)
.ensure_conversation(thread_id, "gateway", &identity.user_id, None)
.await
{
Ok(true) => {}
+2 -2
View File
@@ -2014,7 +2014,7 @@ async fn chat_new_thread_handler(
let session = session_manager.get_or_create_session(&user.user_id).await;
let (thread_id, info) = {
let mut sess = session.lock().await;
let thread = sess.create_thread(Some("gateway"));
let thread = sess.create_thread();
let id = thread.id;
let info = ThreadInfo {
id: thread.id,
@@ -2033,7 +2033,7 @@ async fn chat_new_thread_handler(
// so that the subsequent loadThreads() call from the frontend sees it.
if let Some(ref store) = state.store {
match store
.ensure_conversation(thread_id, "gateway", &user.user_id, None, Some("gateway"))
.ensure_conversation(thread_id, "gateway", &user.user_id, None)
.await
{
Ok(true) => {}
+3 -26
View File
@@ -67,20 +67,19 @@ impl ConversationStore for LibSqlBackend {
channel: &str,
user_id: &str,
thread_id: Option<&str>,
source_channel: Option<&str>,
) -> Result<bool, DatabaseError> {
let conn = self.connect().await?;
let now = fmt_ts(&Utc::now());
let affected = conn
.execute(
r#"
INSERT INTO conversations (id, channel, user_id, thread_id, source_channel, started_at, last_activity)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?6)
INSERT INTO conversations (id, channel, user_id, thread_id, started_at, last_activity)
VALUES (?1, ?2, ?3, ?4, ?5, ?5)
ON CONFLICT (id) DO UPDATE SET last_activity = excluded.last_activity
WHERE conversations.user_id = excluded.user_id
AND conversations.channel = excluded.channel
"#,
params![id.to_string(), channel, user_id, opt_text(thread_id), opt_text(source_channel), now],
params![id.to_string(), channel, user_id, opt_text(thread_id), now],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
@@ -566,28 +565,6 @@ impl ConversationStore for LibSqlBackend {
.map_err(|e| DatabaseError::Query(e.to_string()))?;
Ok(found.is_some())
}
async fn get_conversation_source_channel(
&self,
conversation_id: Uuid,
) -> Result<Option<String>, DatabaseError> {
let conn = self.connect().await?;
let mut rows = conn
.query(
"SELECT source_channel FROM conversations WHERE id = ?1",
params![conversation_id.to_string()],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
match rows
.next()
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
{
Some(row) => Ok(get_opt_text(&row, 0)),
None => Ok(None),
}
}
}
#[cfg(test)]
-8
View File
@@ -785,14 +785,6 @@ CREATE TABLE IF NOT EXISTS api_tokens (
);
CREATE INDEX IF NOT EXISTS idx_api_tokens_user ON api_tokens(user_id);
CREATE INDEX IF NOT EXISTS idx_api_tokens_hash ON api_tokens(token_hash);
"#,
),
(
15,
"conversation_source_channel",
// Add source_channel to conversations for cross-channel approval authorization.
r#"
ALTER TABLE conversations ADD COLUMN source_channel TEXT;
"#,
),
];
-6
View File
@@ -373,7 +373,6 @@ pub trait ConversationStore: Send + Sync {
channel: &str,
user_id: &str,
thread_id: Option<&str>,
source_channel: Option<&str>,
) -> Result<bool, DatabaseError>;
async fn list_conversations_with_preview(
&self,
@@ -432,11 +431,6 @@ pub trait ConversationStore: Send + Sync {
conversation_id: Uuid,
user_id: &str,
) -> Result<bool, DatabaseError>;
/// Get the source_channel for a conversation (the channel that created it).
async fn get_conversation_source_channel(
&self,
conversation_id: Uuid,
) -> Result<Option<String>, DatabaseError>;
}
#[async_trait]
+1 -11
View File
@@ -99,10 +99,9 @@ impl ConversationStore for PgBackend {
channel: &str,
user_id: &str,
thread_id: Option<&str>,
source_channel: Option<&str>,
) -> Result<bool, DatabaseError> {
self.store
.ensure_conversation(id, channel, user_id, thread_id, source_channel)
.ensure_conversation(id, channel, user_id, thread_id)
.await
}
@@ -213,15 +212,6 @@ impl ConversationStore for PgBackend {
.conversation_belongs_to_user(conversation_id, user_id)
.await
}
async fn get_conversation_source_channel(
&self,
conversation_id: Uuid,
) -> Result<Option<String>, DatabaseError> {
self.store
.get_conversation_source_channel(conversation_id)
.await
}
}
// ==================== JobStore ====================
+129 -3
View File
@@ -395,16 +395,49 @@ pub enum RoutineError {
Database { reason: String },
#[error("LLM call failed: {reason}")]
LlmFailed { reason: String },
LlmFailed {
reason: String,
/// Partial token count consumed before the failure (if any).
/// Used to accumulate usage across retry attempts.
partial_tokens: Option<i32>,
/// Whether the underlying LLM error was classified as retryable.
/// Set at the `LlmError` → `RoutineError` conversion site using
/// `crate::llm::retry::is_retryable()`, avoiding fragile substring
/// matching on the stringified reason.
retryable: bool,
},
#[error("Failed to dispatch full job: {reason}")]
JobDispatchFailed { reason: String },
#[error("LLM returned empty content")]
EmptyResponse,
EmptyResponse {
/// Tokens consumed by the call that produced the empty response.
partial_tokens: Option<i32>,
},
#[error("LLM response truncated (finish_reason=length) with no content")]
TruncatedResponse,
TruncatedResponse {
/// Tokens consumed by the call that produced the truncated response.
partial_tokens: Option<i32>,
},
}
impl RoutineError {
/// Whether this error is transient and worth retrying with backoff.
///
/// Retryable: LLM failures where the underlying `LlmError` was classified
/// as retryable by `crate::llm::retry::is_retryable()`, empty responses,
/// and truncated responses.
/// Non-retryable: configuration errors, authorization, resource limits,
/// DB errors, and LLM failures caused by auth/content-policy/context-length.
pub fn is_retryable(&self) -> bool {
match self {
RoutineError::LlmFailed { retryable, .. } => *retryable,
RoutineError::EmptyResponse { .. } | RoutineError::TruncatedResponse { .. } => true,
_ => false,
}
}
}
/// Result type alias for the agent.
@@ -514,6 +547,99 @@ mod tests {
assert!(msg.contains("bad format"), "Should mention reason: {msg}");
}
#[test]
fn routine_error_retryable_classification() {
// Transient errors should be retryable
assert!(
RoutineError::LlmFailed {
reason: "timeout".into(),
partial_tokens: None,
retryable: true,
}
.is_retryable()
);
// Non-retryable LLM error
assert!(
!RoutineError::LlmFailed {
reason: "timeout".into(),
partial_tokens: None,
retryable: false,
}
.is_retryable()
);
assert!(
RoutineError::EmptyResponse {
partial_tokens: None
}
.is_retryable()
);
assert!(
RoutineError::TruncatedResponse {
partial_tokens: None
}
.is_retryable()
);
// Hard failures should NOT be retryable
assert!(
!RoutineError::Disabled {
name: "test".into()
}
.is_retryable()
);
assert!(
!RoutineError::JobDispatchFailed {
reason: "no docker".into()
}
.is_retryable()
);
assert!(
!RoutineError::Database {
reason: "conn refused".into()
}
.is_retryable()
);
assert!(!RoutineError::NotFound { id: Uuid::new_v4() }.is_retryable());
assert!(!RoutineError::NotAuthorized { id: Uuid::new_v4() }.is_retryable());
assert!(
!RoutineError::MaxConcurrent {
name: "test".into()
}
.is_retryable()
);
assert!(
!RoutineError::UnknownTriggerType {
trigger_type: "x".into()
}
.is_retryable()
);
assert!(
!RoutineError::UnknownActionType {
action_type: "x".into()
}
.is_retryable()
);
assert!(
!RoutineError::MissingField {
context: "c".into(),
field: "f".into()
}
.is_retryable()
);
assert!(
!RoutineError::InvalidCron {
reason: "bad".into()
}
.is_retryable()
);
assert!(
!RoutineError::UnknownRunStatus {
status: "bad".into()
}
.is_retryable()
);
}
#[test]
fn top_level_error_from_conversions() {
let config_err = ConfigError::MissingEnvVar("TEST".to_string());
+3 -19
View File
@@ -1581,20 +1581,19 @@ impl Store {
channel: &str,
user_id: &str,
thread_id: Option<&str>,
source_channel: Option<&str>,
) -> Result<bool, DatabaseError> {
let conn = self.conn().await?;
let affected = conn
.execute(
r#"
INSERT INTO conversations (id, channel, user_id, thread_id, source_channel)
VALUES ($1, $2, $3, $4, $5)
INSERT INTO conversations (id, channel, user_id, thread_id)
VALUES ($1, $2, $3, $4)
ON CONFLICT (id) DO UPDATE
SET last_activity = NOW()
WHERE conversations.user_id = EXCLUDED.user_id
AND conversations.channel = EXCLUDED.channel
"#,
&[&id, &channel, &user_id, &thread_id, &source_channel],
&[&id, &channel, &user_id, &thread_id],
)
.await?;
Ok(affected > 0)
@@ -1893,21 +1892,6 @@ impl Store {
Ok(row.is_some())
}
/// Get the source_channel for a conversation.
pub async fn get_conversation_source_channel(
&self,
conversation_id: Uuid,
) -> Result<Option<String>, DatabaseError> {
let conn = self.conn().await?;
let row = conn
.query_opt(
"SELECT source_channel FROM conversations WHERE id = $1",
&[&conversation_id],
)
.await?;
Ok(row.and_then(|r| r.get::<_, Option<String>>(0)))
}
/// Load messages for a conversation with cursor-based pagination.
///
/// Returns `(messages_oldest_first, has_more)`.
+1
View File
@@ -234,6 +234,7 @@ fn is_transient(err: &LlmError) -> bool {
LlmError::RequestFailed { .. }
| LlmError::RateLimited { .. }
| LlmError::InvalidResponse { .. }
| LlmError::EmptyResponse { .. }
| LlmError::SessionExpired { .. }
| LlmError::SessionRenewalFailed { .. }
| LlmError::Http(_)
+3
View File
@@ -17,6 +17,9 @@ pub enum LlmError {
#[error("Invalid response from {provider}: {reason}")]
InvalidResponse { provider: String, reason: String },
#[error("Empty response from {provider}: no content returned")]
EmptyResponse { provider: String },
#[error("Context length exceeded: {used} tokens used, {limit} allowed")]
ContextLengthExceeded { used: usize, limit: usize },
+2 -4
View File
@@ -231,9 +231,8 @@ impl LlmProvider for GithubCopilotProvider {
.choices
.into_iter()
.next()
.ok_or_else(|| LlmError::InvalidResponse {
.ok_or_else(|| LlmError::EmptyResponse {
provider: "github_copilot".to_string(),
reason: "No choices in response".to_string(),
})?;
let (content, _tool_calls) = extract_choice_content(&choice);
@@ -309,9 +308,8 @@ impl LlmProvider for GithubCopilotProvider {
.choices
.into_iter()
.next()
.ok_or_else(|| LlmError::InvalidResponse {
.ok_or_else(|| LlmError::EmptyResponse {
provider: "github_copilot".to_string(),
reason: "No choices in response".to_string(),
})?;
let (content, tool_calls) = extract_choice_content(&choice);
+2 -4
View File
@@ -490,9 +490,8 @@ impl LlmProvider for NearAiChatProvider {
.choices
.into_iter()
.next()
.ok_or_else(|| LlmError::InvalidResponse {
.ok_or_else(|| LlmError::EmptyResponse {
provider: "nearai_chat".to_string(),
reason: "No choices in response".to_string(),
})?;
// Fall back to reasoning_content when content is null (same as
@@ -570,9 +569,8 @@ impl LlmProvider for NearAiChatProvider {
.choices
.into_iter()
.next()
.ok_or_else(|| LlmError::InvalidResponse {
.ok_or_else(|| LlmError::EmptyResponse {
provider: "nearai_chat".to_string(),
reason: "No choices in response".to_string(),
})?;
let tool_calls: Vec<ToolCall> = choice
+1
View File
@@ -48,6 +48,7 @@ pub(crate) fn is_retryable(err: &LlmError) -> bool {
LlmError::RequestFailed { .. }
| LlmError::RateLimited { .. }
| LlmError::InvalidResponse { .. }
| LlmError::EmptyResponse { .. }
| LlmError::SessionRenewalFailed { .. }
| LlmError::Http(_)
| LlmError::Io(_)
-1
View File
@@ -449,7 +449,6 @@ async fn async_main() -> anyhow::Result<()> {
&components.secrets_store,
components.extension_manager.as_ref(),
components.db.as_ref(),
&channel_names,
)
.await;
+1 -1
View File
@@ -278,7 +278,7 @@ impl TenantScope {
thread_id: Option<&str>,
) -> Result<bool, DatabaseError> {
self.inner
.ensure_conversation(id, channel, &self.user_id, thread_id, None)
.ensure_conversation(id, channel, &self.user_id, thread_id)
.await
}
+3 -3
View File
@@ -753,7 +753,7 @@ mod tests {
// ensure_conversation should create the row.
assert!(
db.ensure_conversation(conv_id, "web", "carol", None, Some("web"))
db.ensure_conversation(conv_id, "web", "carol", None)
.await
.expect("ensure first"),
"first ensure_conversation should create the row"
@@ -761,7 +761,7 @@ mod tests {
// Calling again with the same ID should not error.
assert!(
db.ensure_conversation(conv_id, "web", "carol", None, Some("web"))
db.ensure_conversation(conv_id, "web", "carol", None)
.await
.expect("ensure second (idempotent)"),
"second ensure_conversation should touch owned row"
@@ -806,7 +806,7 @@ mod tests {
tokio::time::sleep(std::time::Duration::from_millis(25)).await;
assert!(
!db.ensure_conversation(conv_id, "web", "mallory", None, None)
!db.ensure_conversation(conv_id, "web", "mallory", None)
.await
.expect("foreign ensure should not error"),
"foreign ensure_conversation should report not ensured"
+144 -36
View File
@@ -31,6 +31,7 @@ use std::time::Duration;
use serde::{Deserialize, Serialize};
use tokio::io::{AsyncBufReadExt, BufReader};
#[cfg(not(unix))]
use tokio::process::Command;
use uuid::Uuid;
@@ -340,6 +341,11 @@ impl ClaudeBridgeRuntime {
/// Spawn a `claude` CLI process and stream its output.
///
/// Uses a PTY on Unix so Node.js line-buffers stdout instead of
/// full-buffering (which causes the bridge to hang on non-TTY pipes).
/// Arguments are passed via `execve` (no shell) — injection-safe by
/// construction.
///
/// Returns the session_id if captured from the `system` init message.
async fn run_claude_session(
&self,
@@ -347,47 +353,102 @@ impl ClaudeBridgeRuntime {
resume_session_id: Option<&str>,
extra_env: &std::collections::HashMap<String, String>,
) -> Result<Option<String>, WorkerError> {
let mut cmd = Command::new("claude");
cmd.arg("-p")
.arg(prompt)
.arg("--output-format")
.arg("stream-json")
.arg("--verbose")
.arg("--max-turns")
.arg(self.config.max_turns.to_string())
.arg("--model")
.arg(&self.config.model);
let max_turns_str = self.config.max_turns.to_string();
if let Some(sid) = resume_session_id {
cmd.arg("--resume").arg(sid);
}
// Inject credentials into the child process environment without
// mutating the global process env (which is unsafe in multi-threaded programs).
cmd.envs(extra_env);
cmd.current_dir("/workspace")
.stdout(std::process::Stdio::piped())
.stderr(std::process::Stdio::piped());
let mut child = cmd.spawn().map_err(|e| WorkerError::ExecutionFailed {
reason: format!("failed to spawn claude: {}", e),
})?;
let stdout = child
.stdout
.take()
.ok_or_else(|| WorkerError::ExecutionFailed {
reason: "failed to capture claude stdout".to_string(),
// Spawn with PTY on Unix to fix Node.js stdout buffering.
// All arguments are passed individually via execve — never through
// a shell interpreter. This eliminates shell injection by construction.
#[cfg(unix)]
let (mut child, stdout, stderr) = {
let (pty, pts) = pty_process::open().map_err(|e| WorkerError::ExecutionFailed {
reason: format!("failed to allocate PTY: {}", e),
})?;
let stderr = child
.stderr
.take()
.ok_or_else(|| WorkerError::ExecutionFailed {
reason: "failed to capture claude stderr".to_string(),
let mut cmd = pty_process::Command::new("claude");
cmd = cmd
.arg("-p")
.arg(prompt)
.arg("--output-format")
.arg("stream-json")
.arg("--verbose")
.arg("--max-turns")
.arg(&max_turns_str)
.arg("--model")
.arg(&self.config.model);
if let Some(sid) = resume_session_id {
cmd = cmd.arg("--resume").arg(sid);
}
cmd = cmd.envs(extra_env.iter());
cmd = cmd.current_dir("/workspace");
// Keep stderr on a separate pipe — pty-process attaches the PTY
// to all fds by default, which would merge stderr into the PTY
// stream and break NDJSON parsing.
cmd = cmd.stderr(std::process::Stdio::piped());
let mut child = cmd.spawn(pts).map_err(|e| WorkerError::ExecutionFailed {
reason: format!("failed to spawn claude with PTY: {}", e),
})?;
let stderr = child
.stderr
.take()
.ok_or_else(|| WorkerError::ExecutionFailed {
reason: "failed to capture claude stderr".to_string(),
})?;
// stdout comes from the PTY master, which implements AsyncRead
let stdout: Box<dyn tokio::io::AsyncRead + Unpin + Send> = Box::new(pty);
(child, stdout, stderr)
};
// Non-Unix fallback (Windows CI) — no PTY, direct spawn.
// Claude bridge only runs in Linux Docker containers, so this path
// exists solely for compilation on Windows targets.
#[cfg(not(unix))]
let (mut child, stdout, stderr) = {
let mut cmd = Command::new("claude");
cmd.arg("-p")
.arg(prompt)
.arg("--output-format")
.arg("stream-json")
.arg("--verbose")
.arg("--max-turns")
.arg(&max_turns_str)
.arg("--model")
.arg(&self.config.model);
if let Some(sid) = resume_session_id {
cmd.arg("--resume").arg(sid);
}
cmd.envs(extra_env);
cmd.current_dir("/workspace")
.stdout(std::process::Stdio::piped())
.stderr(std::process::Stdio::piped());
let mut child = cmd.spawn().map_err(|e| WorkerError::ExecutionFailed {
reason: format!("failed to spawn claude: {}", e),
})?;
let stdout_pipe = child
.stdout
.take()
.ok_or_else(|| WorkerError::ExecutionFailed {
reason: "failed to capture claude stdout".to_string(),
})?;
let stderr = child
.stderr
.take()
.ok_or_else(|| WorkerError::ExecutionFailed {
reason: "failed to capture claude stderr".to_string(),
})?;
let stdout: Box<dyn tokio::io::AsyncRead + Unpin + Send> = Box::new(stdout_pipe);
(child, stdout, stderr)
};
// Spawn stderr reader that forwards lines as log events
let client_for_stderr = Arc::clone(&self.client);
let job_id = self.config.job_id;
@@ -1027,4 +1088,51 @@ mod tests {
let copied = copy_dir_recursive(nonexistent, dst.path()).unwrap();
assert_eq!(copied, 0);
}
/// Regression test: arguments are passed individually (not via shell string),
/// so shell metacharacters in prompt/model/session_id are harmless.
#[test]
fn command_args_no_shell_interpretation() {
// Prompt, model, and session_id may contain shell metacharacters from
// user-supplied task descriptions or LLM output. Since we use
// Command::arg() (execve), these are passed as literal strings.
let prompt = "Fix the user's bug; echo $HOME && rm -rf /";
let model = "claude-3-opus-20240229";
let session_id = "'; DROP TABLE jobs; --";
let max_turns = 10u32;
let max_turns_str = max_turns.to_string();
let args: Vec<&str> = vec![
"-p",
prompt,
"--output-format",
"stream-json",
"--verbose",
"--max-turns",
&max_turns_str,
"--model",
model,
"--resume",
session_id,
];
// All values present as literal strings — no shell interpretation
// ["-p", prompt, "--output-format", "stream-json", "--verbose",
// "--max-turns", "10", "--model", model, "--resume", session_id]
assert_eq!(args[1], prompt);
assert_eq!(args[8], model);
assert_eq!(args[10], session_id);
// Shell metacharacters preserved, not expanded
assert!(args[1].contains("$HOME"));
assert!(args[1].contains("&&"));
assert!(args[10].contains("'; DROP TABLE"));
}
/// Verify PTY is available on Unix platforms.
#[cfg(unix)]
#[tokio::test]
async fn pty_opens_successfully() {
let result = pty_process::open();
assert!(result.is_ok(), "PTY allocation should succeed on Unix");
}
}
+148 -4
View File
@@ -391,6 +391,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
worker: self,
rx: tokio::sync::Mutex::new(rx),
consecutive_rate_limits: std::sync::atomic::AtomicUsize::new(0),
has_text_response: std::sync::atomic::AtomicBool::new(false),
};
let config = AgenticLoopConfig {
@@ -1101,6 +1102,15 @@ fn store_fallback_in_metadata(
}
/// Job delegate: implements `LoopDelegate` for the background job context.
/// Whether an LLM error represents a completion-eligible empty response.
///
/// Only `EmptyResponse` (provider returned no choices/content) qualifies.
/// Infrastructure errors (`AuthFailed`, `Http`, `Io`, etc.) never qualify —
/// they must propagate even if prior text output was produced.
fn is_completion_eligible_error(error: &crate::error::LlmError) -> bool {
matches!(error, crate::error::LlmError::EmptyResponse { .. })
}
///
/// Handles: signal channel (stop/ping/user messages), cancellation checks,
/// rate-limit retry, parallel tool execution, DB persistence, SSE broadcasting.
@@ -1109,6 +1119,10 @@ struct JobDelegate<'a> {
rx: tokio::sync::Mutex<&'a mut mpsc::Receiver<WorkerMessage>>,
/// Tracks consecutive rate-limit errors to fail fast instead of burning iterations.
consecutive_rate_limits: std::sync::atomic::AtomicUsize,
/// Whether a substantive (non-empty) text response has been produced.
/// When true, an empty follow-up response is treated as job completion
/// rather than a retry signal (prevents spurious failures in routines).
has_text_response: std::sync::atomic::AtomicBool,
}
impl<'a> JobDelegate<'a> {
@@ -1161,6 +1175,53 @@ impl<'a> JobDelegate<'a> {
finish_reason: crate::llm::FinishReason::Stop,
})
}
/// Mark the job as completed, logging a warning on failure.
async fn mark_completed_or_warn(&self, context: &str) {
if let Err(e) = self.worker.mark_completed().await {
tracing::warn!(
job_id = %self.worker.job_id,
error = %e,
"Failed to mark job completed ({context})"
);
}
}
/// If a substantive text response was already produced and the error
/// indicates the LLM simply returned nothing, treat it as successful
/// completion rather than a fatal failure.
///
/// Only swallows `EmptyResponse` — infrastructure errors (`AuthFailed`,
/// `ContextLengthExceeded`, `Http`, `Io`, etc.) always propagate.
///
/// Returns `Some(empty RespondOutput)` when the error should be swallowed,
/// `None` when it should propagate normally.
async fn try_complete_on_error(
&self,
context: &str,
error: &crate::error::LlmError,
) -> Option<crate::llm::RespondOutput> {
if !is_completion_eligible_error(error) {
return None;
}
if !self
.has_text_response
.load(std::sync::atomic::Ordering::Relaxed)
{
return None;
}
tracing::info!(
job_id = %self.worker.job_id,
error = %error,
"{context} empty response after text output — treating as completion"
);
self.mark_completed_or_warn(context).await;
Some(crate::llm::RespondOutput {
result: RespondResult::Text(String::new()),
usage: crate::llm::TokenUsage::default(),
finish_reason: crate::llm::FinishReason::Stop,
})
}
}
#[async_trait]
@@ -1291,7 +1352,12 @@ impl<'a> LoopDelegate for JobDelegate<'a> {
Err(crate::error::LlmError::RateLimited { retry_after, .. }) => {
return self.handle_rate_limit(retry_after, "tool selection").await;
}
Err(e) => return Err(e.into()),
Err(e) => {
if let Some(output) = self.try_complete_on_error("select_tools", &e).await {
return Ok(output);
}
return Err(e.into());
}
};
// Fall back to respond_with_tools
@@ -1321,7 +1387,12 @@ impl<'a> LoopDelegate for JobDelegate<'a> {
self.handle_rate_limit(retry_after, "respond_with_tools")
.await
}
Err(e) => Err(e.into()),
Err(e) => {
if let Some(output) = self.try_complete_on_error("respond_with_tools", &e).await {
return Ok(output);
}
Err(e.into())
}
}
}
@@ -1330,9 +1401,22 @@ impl<'a> LoopDelegate for JobDelegate<'a> {
text: &str,
reason_ctx: &mut ReasoningContext,
) -> TextAction {
// Empty text from rate-limit backoff retry — skip processing and let the
// loop proceed to the next iteration which will re-call the LLM.
// Empty text after a substantive response means the LLM has finished.
// Treat as successful completion rather than continuing the loop (which
// would produce "Response contained no message or tool call (empty)").
if text.is_empty() {
if self
.has_text_response
.load(std::sync::atomic::Ordering::Relaxed)
{
tracing::debug!(
job_id = %self.worker.job_id,
"Empty response after text output — treating as completion"
);
self.mark_completed_or_warn("empty text response").await;
return TextAction::Return(LoopOutcome::Response(String::new()));
}
// No prior text response — this is likely a rate-limit backoff retry.
return TextAction::Continue;
}
@@ -1348,6 +1432,10 @@ impl<'a> LoopDelegate for JobDelegate<'a> {
return TextAction::Return(LoopOutcome::Response(text.to_string()));
}
// Track that a substantive response has been produced.
self.has_text_response
.store(true, std::sync::atomic::Ordering::Relaxed);
// Add assistant response to context
reason_ctx.messages.push(ChatMessage::assistant(text));
@@ -2285,4 +2373,60 @@ mod tests {
assert_eq!(telegram[0].0, "owner-scope");
assert_eq!(telegram[0].1.content, "hello from routine");
}
/// Regression test: only `EmptyResponse` errors are eligible for
/// completion-swallowing. Infrastructure errors must always propagate.
#[test]
fn is_completion_eligible_only_matches_empty_response() {
use crate::error::LlmError;
// EmptyResponse is eligible
assert!(super::is_completion_eligible_error(
&LlmError::EmptyResponse {
provider: "test".to_string(),
}
));
// All other variants are NOT eligible
assert!(!super::is_completion_eligible_error(
&LlmError::InvalidResponse {
provider: "test".to_string(),
reason: "parse error".to_string(),
}
));
assert!(!super::is_completion_eligible_error(
&LlmError::AuthFailed {
provider: "test".to_string(),
}
));
assert!(!super::is_completion_eligible_error(
&LlmError::ContextLengthExceeded {
used: 100_000,
limit: 50_000,
}
));
assert!(!super::is_completion_eligible_error(
&LlmError::ModelNotAvailable {
provider: "test".to_string(),
model: "gpt-4".to_string(),
}
));
assert!(!super::is_completion_eligible_error(
&LlmError::RequestFailed {
provider: "test".to_string(),
reason: "timeout".to_string(),
}
));
assert!(!super::is_completion_eligible_error(
&LlmError::SessionExpired {
provider: "test".to_string(),
}
));
assert!(!super::is_completion_eligible_error(
&LlmError::SessionRenewalFailed {
provider: "test".to_string(),
reason: "timeout".to_string(),
}
));
}
}
+1 -7
View File
@@ -48,13 +48,7 @@ mod tests {
let store = rig.database();
assert!(
store
.ensure_conversation(
foreign_thread_id,
"gateway",
"victim-user",
None,
Some("gateway")
)
.ensure_conversation(foreign_thread_id, "gateway", "victim-user", None)
.await
.expect("failed to create victim conversation"),
"test setup failed: victim conversation was not created"