Compare commits

...
Author SHA1 Message Date
Claude f61a0b747f fix(clippy): use any() instead of find().is_none() in user stats test
https://claude.ai/code/session_01Nm95eCjdrwxDjwkHTRieZs
2026-03-28 19:21:15 +00:00
ZakiandClaude 65c374f53c fix(routines): broaden strip_html_tags to cover all HTML forms
The whitelist-based regex missed self-closing tags without whitespace
(<br/>, <img/>), HTML comments (<!--...-->), SVG/MathML tags, and
custom elements (<custom-element>). This weakened the HTML stripping
guarantee for untrusted routine/job summaries in notifications.

Changes:
- Add separate regex for HTML comments (<!--...-->)
- Add SVG tags (svg, path, circle, etc.) and MathML tags (math, mrow,
  etc.) to the known tag list
- Add regex for custom elements (tags containing hyphens per web
  components spec)
- Fix self-closing tag matching to handle <br/> without whitespace
  by making the whitespace before /> optional
- Add regression tests for all four cases plus generics preservation
- Fix pre-existing compilation error in tunnel/mod.rs test helpers
  (missing GatewayConfig fields)

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-28 19:16:32 +00:00
Claude faf3012a36 fix: remove .expect() from strip_html_tags to pass no-panics check
Use Option<Regex> with graceful fallback instead of panicking on
regex compilation failure.

https://claude.ai/code/session_01CsP5wMZ2evEMHgGghjAfR1
2026-03-28 19:16:03 +00:00
ZakiandClaude ce193ff2d2 fix(routines): set run.job_id in-memory and revert Cargo.lock drift
Address PR #1470 review feedback:

- Set run.job_id = Some(job_id) after link_routine_run_to_job succeeds
  so send_notification reads the correct value instead of always None.
- Revert Cargo.lock to staging baseline: the PR had accumulated
  unrelated dependency changes (openssl, native-tls, crossterm 0.28.1
  downgrade, foreign-types, vcpkg) from a dirty lockfile resolution.

[skip-regression-check]

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-28 19:16:03 +00:00
ZakiandClaude 7881998d93 fix: preserve non-HTML angle brackets in sanitize_summary
The previous strip_html_tags() implementation blindly removed all content
between angle brackets, mangling legitimate text like Vec<String>, shell
redirects (cat < input.txt), and comparison operators in LLM/error output.

Replace the naive char-by-char scanner with a regex that only matches
known HTML tag names (div, script, img, a, b, etc.), preserving generic
angle-bracket content that appears in code snippets and error messages.

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-28 19:16:03 +00:00
Claude d967e620cc fix: resolve mutability mismatch in execute_full_job call after rebase
[skip-regression-check]

https://claude.ai/code/session_01ABGWibdKVQ3b6pEKtxPPkM
2026-03-28 19:16:03 +00:00
ZakiandClaude 7b884c4a24 fix(routines): propagate job_id to notification metadata (#1321)
Update run.job_id after execute_full_job() links the routine run to the
job, ensuring send_notification receives the actual job ID instead of None
on the normal completion path.

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-28 19:16:03 +00:00
Claude d45d5977a0 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:03 +00:00
Claude d29811db6b style: run cargo fmt to fix formatting
https://claude.ai/code/session_01Nv2TJ3so5WQqpRUcirAhT3
2026-03-28 19:16:03 +00:00
ZakiandClaude 25dcbf78b8 fix(routines): normalize notification summaries with truncation and metadata (#1321)
- Capitalize status labels in notifications (ok -> Completed, attention -> Needs attention)
- Sanitize and truncate long summaries to 500 chars with UTF-8-safe ellipsis
- Include job_id in notification metadata for full-job routines
- Move sanitize_summary/strip_html_tags out of #[cfg(test)] for production use
- Add regression tests for truncation, status labels, job_id metadata

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-28 19:16:03 +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
AchieveandGitHub 9bb19a98f7 fix(web): redact database error details from API responses (#1711) 2026-03-28 15:13:28 +01:00
AchieveandGitHub 0b33ca9926 fix(oauth): tighten legacy state validation and fallback handling (#1701)
* fix(oauth): tighten legacy state validation and fallback handling

* style: fix formatting

* refactor: separate validation checks for clearer error messages
2026-03-28 15:10:39 +01:00
AchieveandGitHub 9ba10eac35 fix(db): add tracing warn for naive timestamp fallback and improve parse_timestamp tests (#1700)
* fix(db): add tracing warn for naive timestamp fallback and improve parse_timestamp tests

* style: fix formatting
2026-03-28 15:08:25 +01:00
AchieveandGitHub 27e8d6f8dd fix(wasm): use typed WASM schema as advertised schema when available (#1699) 2026-03-28 15:07:16 +01:00
Henry ParkandGitHub f49f368355 Clean up extension credentials on uninstall (#1718)
* Clean up extension credentials on uninstall

* Address PR review feedback

* Cover channel webhook secrets on uninstall

* Harden tool secret cleanup detection
2026-03-28 14:46:45 +01:00
21 changed files with 2182 additions and 165 deletions
Generated
+141 -11
View File
@@ -1510,7 +1510,7 @@ version = "1.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "980c2afde4af43d6a05c5be738f9eae595cff86dce1f38f88b95058a98c027f3"
dependencies = [
"crossterm",
"crossterm 0.29.0",
]
[[package]]
@@ -1731,7 +1731,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "04a63daf06a168535c74ab97cdba3ed4fa5d4f32cb36e437dcceb83d66854b7c"
dependencies = [
"crokey-proc_macros",
"crossterm",
"crossterm 0.29.0",
"once_cell",
"serde",
"strict",
@@ -1743,7 +1743,7 @@ version = "1.4.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "847f11a14855fc490bd5d059821895c53e77eeb3c2b73ee3dded7ce77c93b231"
dependencies = [
"crossterm",
"crossterm 0.29.0",
"proc-macro2",
"quote",
"strict",
@@ -1817,6 +1817,22 @@ version = "0.8.21"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d0a5c400df2834b80a4c3327b3aad3a4c4cd4de0629063962b03235697506a28"
[[package]]
name = "crossterm"
version = "0.28.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "829d955a0bb380ef178a640b91779e3987da38c9aea133b20614cfed8cdea9c6"
dependencies = [
"bitflags 2.11.0",
"crossterm_winapi",
"mio",
"parking_lot",
"rustix 0.38.44",
"signal-hook",
"signal-hook-mio",
"winapi",
]
[[package]]
name = "crossterm"
version = "0.29.0"
@@ -2476,6 +2492,21 @@ version = "0.2.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "77ce24cb58228fbb8aa041425bb1050850ac19177686ea6e0f41a70416f56fdb"
[[package]]
name = "foreign-types"
version = "0.3.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f6f339eb8adc052cd2ca78910fda869aefa38d22d5cb648e6485e4d3fc06f3b1"
dependencies = [
"foreign-types-shared",
]
[[package]]
name = "foreign-types-shared"
version = "0.1.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "00b0228411908ca8685dba7fc2cdd70ec9990a6e753e89b6ac91a84c40fbaf4b"
[[package]]
name = "form_urlencoded"
version = "1.2.2"
@@ -3118,7 +3149,6 @@ dependencies = [
"tokio",
"tokio-rustls 0.26.4",
"tower-service",
"webpki-roots 1.0.6",
]
[[package]]
@@ -3133,6 +3163,22 @@ dependencies = [
"tokio-io-timeout",
]
[[package]]
name = "hyper-tls"
version = "0.6.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "70206fc6890eaca9fde8a0bf71caa2ddfc9fe045ac9e5c70df101a7dbde866e0"
dependencies = [
"bytes",
"http-body-util",
"hyper 1.8.1",
"hyper-util",
"native-tls",
"tokio",
"tokio-native-tls",
"tower-service",
]
[[package]]
name = "hyper-util"
version = "0.1.20"
@@ -3150,7 +3196,7 @@ dependencies = [
"libc",
"percent-encoding",
"pin-project-lite",
"socket2 0.5.10",
"socket2 0.6.3",
"system-configuration",
"tokio",
"tower-service",
@@ -3410,7 +3456,7 @@ dependencies = [
"clap_complete",
"criterion",
"cron",
"crossterm",
"crossterm 0.28.1",
"deadpool-postgres",
"dirs 6.0.0",
"dotenvy",
@@ -3439,6 +3485,7 @@ dependencies = [
"pgvector",
"postgres-types",
"pretty_assertions",
"pty-process",
"rand 0.8.5",
"readabilityrs",
"refinery",
@@ -3524,7 +3571,7 @@ checksum = "3640c1c38b8e4e43584d8df18be5fc6b0aa314ce6ebf51b53313d4306cca8e46"
dependencies = [
"hermit-abi",
"libc",
"windows-sys 0.59.0",
"windows-sys 0.61.2",
]
[[package]]
@@ -4088,6 +4135,23 @@ dependencies = [
"rand 0.8.5",
]
[[package]]
name = "native-tls"
version = "0.2.18"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "465500e14ea162429d264d44189adc38b199b62b1c21eea9f69e4b73cb03bbf2"
dependencies = [
"libc",
"log",
"openssl",
"openssl-probe 0.2.1",
"openssl-sys",
"schannel",
"security-framework 3.7.0",
"security-framework-sys",
"tempfile",
]
[[package]]
name = "new_debug_unreachable"
version = "1.0.6"
@@ -4310,6 +4374,32 @@ dependencies = [
"pathdiff",
]
[[package]]
name = "openssl"
version = "0.10.76"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "951c002c75e16ea2c65b8c7e4d3d51d5530d8dfa7d060b4776828c88cfb18ecf"
dependencies = [
"bitflags 2.11.0",
"cfg-if",
"foreign-types",
"libc",
"once_cell",
"openssl-macros",
"openssl-sys",
]
[[package]]
name = "openssl-macros"
version = "0.1.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a948666b637a0f465e8564c73e89d4dde00d72d4d473cc972f390fc3dcee7d9c"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.117",
]
[[package]]
name = "openssl-probe"
version = "0.1.6"
@@ -4322,6 +4412,18 @@ version = "0.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7c87def4c32ab89d880effc9e097653c8da5d6ef28e6b539d313baaacfbafcbe"
[[package]]
name = "openssl-sys"
version = "0.9.112"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "57d55af3b3e226502be1526dfdba67ab0e9c96fc293004e79576b2b9edb0dbdb"
dependencies = [
"cc",
"libc",
"pkg-config",
"vcpkg",
]
[[package]]
name = "option-ext"
version = "0.2.0"
@@ -4906,6 +5008,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 +5042,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 +5079,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]]
@@ -5301,11 +5413,13 @@ dependencies = [
"http-body-util",
"hyper 1.8.1",
"hyper-rustls 0.27.7",
"hyper-tls",
"hyper-util",
"js-sys",
"log",
"mime",
"mime_guess",
"native-tls",
"percent-encoding",
"pin-project-lite",
"quinn",
@@ -5317,6 +5431,7 @@ dependencies = [
"serde_urlencoded",
"sync_wrapper 1.0.2",
"tokio",
"tokio-native-tls",
"tokio-rustls 0.26.4",
"tokio-util",
"tower 0.5.3",
@@ -5327,7 +5442,6 @@ dependencies = [
"wasm-bindgen-futures",
"wasm-streams",
"web-sys",
"webpki-roots 1.0.6",
]
[[package]]
@@ -6660,6 +6774,16 @@ dependencies = [
"syn 2.0.117",
]
[[package]]
name = "tokio-native-tls"
version = "0.3.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "bbae76ab933c85776efabc971569dd6119c580d8f5d448769dec1764bf796ef2"
dependencies = [
"native-tls",
"tokio",
]
[[package]]
name = "tokio-postgres"
version = "0.7.16"
@@ -7343,6 +7467,12 @@ version = "0.1.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "ba73ea9cf16a25df0c8caa16c51acb937d5712a8429db78a3ee29d5dcacd3a65"
[[package]]
name = "vcpkg"
version = "0.2.15"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "accd4ea62f7bb7a82fe23066fb0957d48ef677f6eeb8215f372f52e48bb32426"
[[package]]
name = "version_check"
version = "0.9.5"
+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
}
}
+326 -25
View File
@@ -715,6 +715,7 @@ impl RoutineEngine {
status,
Some(summary),
thread_id.as_deref(),
run.job_id,
)
.await;
@@ -1085,7 +1086,7 @@ struct EngineContext {
}
/// Execute a routine run. Handles both lightweight and full_job modes.
async fn execute_routine(ctx: EngineContext, routine: Routine, run: RoutineRun) {
async fn execute_routine(ctx: EngineContext, routine: Routine, mut run: RoutineRun) {
// Increment running count (atomic: survives panics in the execution below)
ctx.running_count.fetch_add(1, Ordering::Relaxed);
@@ -1118,7 +1119,7 @@ async fn execute_routine(ctx: EngineContext, routine: Routine, run: RoutineRun)
description,
max_iterations: *max_iterations,
};
execute_full_job(&ctx, &routine, &run, &execution).await
execute_full_job(&ctx, &routine, &mut run, &execution).await
}
};
@@ -1218,6 +1219,7 @@ async fn execute_routine(ctx: EngineContext, routine: Routine, run: RoutineRun)
status,
summary.as_deref(),
thread_id.as_deref(),
run.job_id,
)
.await;
}
@@ -1253,7 +1255,7 @@ struct FullJobExecutionConfig<'a> {
async fn execute_full_job(
ctx: &EngineContext,
routine: &Routine,
run: &RoutineRun,
run: &mut RoutineRun,
execution: &FullJobExecutionConfig<'_>,
) -> Result<(RunStatus, Option<String>, Option<i32>), RoutineError> {
match ctx.sandbox_readiness {
@@ -1315,6 +1317,9 @@ async fn execute_full_job(
reason: format!("failed to link run to job: {e}"),
})?;
// Keep the in-memory struct in sync so send_notification can read run.job_id.
run.job_id = Some(job_id);
tracing::info!(
routine = %routine.name,
job_id = %job_id,
@@ -1827,7 +1832,18 @@ async fn execute_routine_tool(
Ok(result_str)
}
/// Human-readable label for a run status, suitable for user-facing notifications.
fn status_display_label(status: RunStatus) -> &'static str {
match status {
RunStatus::Ok => "Completed",
RunStatus::Attention => "Needs attention",
RunStatus::Failed => "Failed",
RunStatus::Running => "Running",
}
}
/// 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,
@@ -1836,6 +1852,7 @@ async fn send_notification(
status: RunStatus,
summary: Option<&str>,
thread_id: Option<&str>,
job_id: Option<Uuid>,
) {
let should_notify = match status {
RunStatus::Ok => notify.on_success,
@@ -1855,23 +1872,37 @@ async fn send_notification(
RunStatus::Running => "",
};
let label = status_display_label(status);
let message = match summary {
Some(s) => format!("{} *Routine '{}'*: {}\n\n{}", icon, routine_name, status, s),
None => format!("{} *Routine '{}'*: {}", icon, routine_name, status),
Some(s) => {
let sanitized = sanitize_summary(s);
format!(
"{} *Routine '{}'*: {}\n\n{}",
icon, routine_name, label, sanitized
)
}
None => format!("{} *Routine '{}'*: {}", icon, routine_name, label),
};
let mut metadata = serde_json::json!({
"source": "routine",
"routine_name": routine_name,
"status": status.to_string(),
"owner_id": owner_id,
"notify_user": notify.user,
"notify_channel": notify.channel,
});
if let Some(jid) = job_id {
metadata["job_id"] = serde_json::json!(jid.to_string());
}
let response = OutgoingResponse {
content: message,
thread_id: thread_id.map(String::from),
attachments: Vec::new(),
metadata: serde_json::json!({
"source": "routine",
"routine_name": routine_name,
"status": status.to_string(),
"owner_id": owner_id,
"notify_user": notify.user,
"notify_channel": notify.channel,
}),
metadata,
};
if let Err(e) = tx.send(response).await {
@@ -1934,7 +1965,6 @@ fn truncate(s: &str, max: usize) -> String {
/// 2. Strip HTML tags to prevent injection in web-rendered notifications
/// 3. Collapse multiple whitespace/newlines to single spaces for cleaner output
/// 4. Truncate to 500 chars to prevent oversized notifications
#[cfg(test)]
fn sanitize_summary(s: &str) -> String {
// Strip control characters (keep newline for now, collapse later)
let no_control: String = s
@@ -1961,19 +1991,59 @@ fn sanitize_summary(s: &str) -> String {
}
}
/// Remove HTML/XML tags from a string.
#[cfg(test)]
/// Remove actual HTML tags from a string while preserving non-HTML angle brackets.
///
/// Only strips patterns that look like real HTML/XML tags (e.g. `<div>`, `</p>`,
/// `<img src=...>`), not generic angle-bracket content like `Vec<String>`,
/// `cat < input.txt`, or comparison operators.
///
/// Also strips HTML comments (`<!--...-->`), SVG/MathML tags, and custom elements
/// (tags containing hyphens like `<custom-element>`).
fn strip_html_tags(s: &str) -> String {
let mut result = String::with_capacity(s.len());
let mut in_tag = false;
for c in s.chars() {
match c {
'<' => in_tag = true,
'>' if in_tag => in_tag = false,
_ if !in_tag => result.push(c),
_ => {}
}
use std::sync::LazyLock;
// HTML comment pattern: <!--...-->
static COMMENT_RE: LazyLock<Option<Regex>> =
LazyLock::new(|| Regex::new(r"<!--[\s\S]*?-->").ok());
// Known HTML/SVG/MathML tag names. Includes SVG tags (svg, path, circle, etc.)
// and MathML tags (math, mrow, etc.) that can carry event handlers.
static HTML_TAG_RE: LazyLock<Option<Regex>> = LazyLock::new(|| {
let tags = "a|abbr|address|area|article|aside|audio|b|base|bdi|bdo|blockquote|\
body|br|button|canvas|caption|cite|code|col|colgroup|data|datalist|dd|del|\
details|dfn|dialog|div|dl|dt|em|embed|fieldset|figcaption|figure|footer|\
form|h[1-6]|head|header|hgroup|hr|html|i|iframe|img|input|ins|kbd|label|\
legend|li|link|main|map|mark|meta|meter|nav|noscript|object|ol|optgroup|\
option|output|p|param|picture|pre|progress|q|rp|rt|ruby|s|samp|script|\
section|select|slot|small|source|span|strong|style|sub|summary|sup|table|\
tbody|td|template|textarea|tfoot|th|thead|time|title|tr|track|u|ul|var|\
video|wbr|\
svg|g|path|circle|ellipse|line|polyline|polygon|rect|text|tspan|defs|\
clippath|mask|pattern|image|use|symbol|marker|lineargradient|\
radialgradient|stop|filter|foreignobject|animate|animatetransform|\
math|mrow|mi|mo|mn|ms|mtext|mfrac|msqrt|mroot|msub|msup|msubsup|\
munder|mover|munderover|mtable|mtr|mtd|mspace|mpadded|mfenced|menclose";
// Handles: <tag>, </tag>, <tag/>, <tag />, <tag attr="val">, <tag attr="val"/>
Regex::new(&format!(r"(?i)</?(?:{})(?:\s[^>]*)?\s*/?>", tags)).ok()
});
// Custom elements: tags containing a hyphen (web components spec requires it).
// E.g. <custom-element>, <my-widget foo="bar">, </x-foo>
static CUSTOM_ELEMENT_RE: LazyLock<Option<Regex>> =
LazyLock::new(|| Regex::new(r"(?i)</?\w+-[\w-]*(?:\s[^>]*)?\s*/?>").ok());
let mut result = s.to_string();
if let Some(re) = COMMENT_RE.as_ref() {
result = re.replace_all(&result, "").into_owned();
}
if let Some(re) = HTML_TAG_RE.as_ref() {
result = re.replace_all(&result, "").into_owned();
}
if let Some(re) = CUSTOM_ELEMENT_RE.as_ref() {
result = re.replace_all(&result, "").into_owned();
}
result
}
@@ -2548,6 +2618,33 @@ mod tests {
assert_eq!(sanitize_summary("<img src=x onerror=alert(1)>"), "");
}
#[test]
fn test_sanitize_summary_preserves_non_html_angle_brackets() {
use super::sanitize_summary;
// Rust/Java generics must pass through unchanged
assert_eq!(
sanitize_summary("expected Vec<String>"),
"expected Vec<String>"
);
assert_eq!(
sanitize_summary("HashMap<String, Vec<u8>>"),
"HashMap<String, Vec<u8>>"
);
// Shell redirects must pass through unchanged
assert_eq!(sanitize_summary("cat < input.txt"), "cat < input.txt");
// Comparison operators must pass through unchanged
assert_eq!(sanitize_summary("x < 10 && y > 20"), "x < 10 && y > 20");
// Mixed: real HTML stripped but generics preserved
assert_eq!(
sanitize_summary("Error in Vec<String>: <b>failed</b>"),
"Error in Vec<String>: failed"
);
}
#[test]
fn test_sanitize_summary_multibyte_truncation() {
use super::sanitize_summary;
@@ -2558,4 +2655,208 @@ mod tests {
assert!(result.len() <= 503);
assert!(result.ends_with("..."));
}
#[test]
fn test_sanitize_summary_truncates_long_text() {
use super::sanitize_summary;
let short = "This is a short summary.";
assert_eq!(sanitize_summary(short), short);
let long = "x".repeat(600);
let result = sanitize_summary(&long);
assert!(
result.len() <= 503,
"Truncated summary should be at most 503 bytes (500 + '...')"
);
assert!(
result.ends_with("..."),
"Truncated summary should end with ellipsis"
);
}
#[test]
fn test_sanitize_summary_strips_all_html_forms() {
use super::sanitize_summary;
// Self-closing tags without whitespace: <br/>, <img/>
assert_eq!(sanitize_summary("line1<br/>line2"), "line1line2");
assert_eq!(sanitize_summary("text<img/>more"), "textmore");
assert_eq!(sanitize_summary("text<br />more"), "textmore");
// HTML comments
assert_eq!(sanitize_summary("before<!--x-->after"), "beforeafter");
assert_eq!(sanitize_summary("a<!-- multi\nline -->b"), "ab");
// SVG tags (can carry event handlers)
assert_eq!(
sanitize_summary("<svg onload=alert(1)>payload</svg>"),
"payload"
);
assert_eq!(sanitize_summary("<svg><circle r=10/></svg>"), "");
// MathML tags
assert_eq!(sanitize_summary("<math><mrow>x</mrow></math>"), "x");
// Custom elements (web components with hyphens)
assert_eq!(
sanitize_summary("before<custom-element>inner</custom-element>after"),
"beforeinnerafter"
);
assert_eq!(
sanitize_summary("<my-widget foo=\"bar\">content</my-widget>"),
"content"
);
// Generics must still be preserved
assert_eq!(
sanitize_summary("expected Vec<String>"),
"expected Vec<String>"
);
}
#[test]
fn test_status_display_label_readable() {
use super::status_display_label;
assert_eq!(status_display_label(RunStatus::Ok), "Completed");
assert_eq!(status_display_label(RunStatus::Failed), "Failed");
assert_eq!(
status_display_label(RunStatus::Attention),
"Needs attention"
);
assert_eq!(status_display_label(RunStatus::Running), "Running");
}
#[tokio::test]
async fn test_notification_message_uses_readable_status() {
use tokio::sync::mpsc;
let (tx, mut rx) = mpsc::channel(1);
let notify = NotifyConfig {
on_success: true,
on_failure: true,
on_attention: true,
..Default::default()
};
super::send_notification(
&tx,
&notify,
"user-1",
"my-routine",
RunStatus::Ok,
Some("All good"),
None,
None,
)
.await;
let msg = rx.recv().await.expect("should receive notification");
assert!(
msg.content.contains("Completed"),
"Notification should use readable label 'Completed', got: {}",
msg.content
);
assert!(
!msg.content.contains(": ok"),
"Notification should not contain raw lowercase status"
);
}
#[tokio::test]
async fn test_notification_includes_job_id_in_metadata() {
use tokio::sync::mpsc;
let (tx, mut rx) = mpsc::channel(1);
let notify = NotifyConfig {
on_failure: true,
..Default::default()
};
let job_id = uuid::Uuid::new_v4();
super::send_notification(
&tx,
&notify,
"user-1",
"my-routine",
RunStatus::Failed,
Some("something broke"),
None,
Some(job_id),
)
.await;
let msg = rx.recv().await.expect("should receive notification");
let meta_job_id = msg.metadata["job_id"]
.as_str()
.expect("metadata should contain job_id");
assert_eq!(meta_job_id, job_id.to_string());
}
#[tokio::test]
async fn test_notification_omits_job_id_when_none() {
use tokio::sync::mpsc;
let (tx, mut rx) = mpsc::channel(1);
let notify = NotifyConfig {
on_success: true,
..Default::default()
};
super::send_notification(
&tx,
&notify,
"user-1",
"my-routine",
RunStatus::Ok,
Some("done"),
None,
None,
)
.await;
let msg = rx.recv().await.expect("should receive notification");
assert!(
msg.metadata.get("job_id").is_none(),
"metadata should not contain job_id when None"
);
}
#[tokio::test]
async fn test_notification_truncates_long_summary() {
use tokio::sync::mpsc;
let (tx, mut rx) = mpsc::channel(1);
let notify = NotifyConfig {
on_failure: true,
..Default::default()
};
let long_summary = "z".repeat(1000);
super::send_notification(
&tx,
&notify,
"user-1",
"my-routine",
RunStatus::Failed,
Some(&long_summary),
None,
None,
)
.await;
let msg = rx.recv().await.expect("should receive notification");
// The sanitized summary should be truncated to ~500 chars + "..."
// The full message includes icon + routine name + label, so just check
// it doesn't contain the full 1000-char string.
assert!(
!msg.content.contains(&long_summary),
"Notification should truncate long summaries"
);
assert!(
msg.content.contains("..."),
"Truncated notification should contain ellipsis"
);
}
}
+30 -32
View File
@@ -15,6 +15,14 @@ use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*;
fn db_error(context: &str, e: impl std::fmt::Display) -> (StatusCode, String) {
tracing::error!(%e, context, "Database error in jobs handler");
(
StatusCode::INTERNAL_SERVER_ERROR,
"Internal database error".to_string(),
)
}
pub async fn jobs_list_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
@@ -213,10 +221,7 @@ pub async fn jobs_detail_handler(
}
Ok(None) => {}
Err(e) => {
return Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
));
return Err(db_error("jobs_handler", e));
}
}
@@ -257,10 +262,7 @@ pub async fn jobs_detail_handler(
}))
}
Ok(None) => Err((StatusCode::NOT_FOUND, "Job not found".to_string())),
Err(e) => Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
)),
Err(e) => Err(db_error("jobs_handler", e)),
}
}
@@ -304,10 +306,7 @@ pub async fn jobs_cancel_handler(
}
Ok(None) => {}
Err(e) => {
return Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
));
return Err(db_error("jobs_handler", e));
}
}
}
@@ -350,10 +349,7 @@ pub async fn jobs_cancel_handler(
}
Ok(None) => {}
Err(e) => {
return Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
));
return Err(db_error("jobs_handler", e));
}
}
}
@@ -471,10 +467,7 @@ pub async fn jobs_restart_handler(
}
Ok(None) => {}
Err(e) => {
return Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
));
return Err(db_error("jobs_handler", e));
}
}
@@ -530,10 +523,7 @@ pub async fn jobs_restart_handler(
})))
}
Ok(None) => Err((StatusCode::NOT_FOUND, "Job not found".to_string())),
Err(e) => Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
)),
Err(e) => Err(db_error("jobs_handler", e)),
}
}
@@ -609,10 +599,7 @@ pub async fn jobs_prompt_handler(
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
Err(e) => {
return Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
));
return Err(db_error("jobs_handler", e));
}
}
}
@@ -667,10 +654,7 @@ pub async fn jobs_events_handler(
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
Err(e) => {
return Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
));
return Err(db_error("jobs_handler", e));
}
}
@@ -823,3 +807,17 @@ pub async fn job_files_read_handler(
content,
}))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_db_error_does_not_leak_details() {
let (status, body) = db_error("test_context", "relation \"jobs\" does not exist");
assert_eq!(status, StatusCode::INTERNAL_SERVER_ERROR);
assert_eq!(body, "Internal database error");
assert!(!body.contains("relation"));
assert!(!body.contains("does not exist"));
}
}
+122 -4
View File
@@ -569,6 +569,42 @@ pub async fn sweep_expired_flows(registry: &PendingOAuthRegistry) {
const HOSTED_STATE_PREFIX: &str = "ic2";
const HOSTED_STATE_CHECKSUM_BYTES: usize = 12;
/// Maximum length for a legacy flow ID or instance name.
const LEGACY_STATE_MAX_LEN: usize = 128;
/// Minimum length for a legacy flow ID.
const LEGACY_STATE_MIN_LEN: usize = 8;
/// Validate that a legacy state component (flow_id or instance_name) contains
/// only safe characters: alphanumeric, dash, underscore.
fn is_valid_legacy_state_component(s: &str) -> bool {
!s.is_empty()
&& s.len() <= LEGACY_STATE_MAX_LEN
&& s.bytes()
.all(|b| b.is_ascii_alphanumeric() || b == b'-' || b == b'_')
}
fn validate_legacy_flow_id(flow_id: &str) -> Result<(), String> {
if flow_id.len() < LEGACY_STATE_MIN_LEN {
return Err(format!(
"Legacy OAuth flow_id too short ({} chars, minimum {LEGACY_STATE_MIN_LEN})",
flow_id.len()
));
}
if flow_id.len() > LEGACY_STATE_MAX_LEN {
return Err(format!(
"Legacy OAuth flow_id too long ({} chars, maximum {LEGACY_STATE_MAX_LEN})",
flow_id.len()
));
}
if !flow_id
.bytes()
.all(|b| b.is_ascii_alphanumeric() || b == b'-' || b == b'_')
{
return Err("Legacy OAuth flow_id contains invalid characters".to_string());
}
Ok(())
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct DecodedHostedOAuthState {
pub flow_id: String,
@@ -653,6 +689,17 @@ pub fn decode_hosted_oauth_state(state: &str) -> Result<DecodedHostedOAuthState,
if flow_id.is_empty() {
return Err("Hosted OAuth legacy state is missing flow_id".to_string());
}
validate_legacy_flow_id(flow_id)?;
if !instance_name.is_empty() && !is_valid_legacy_state_component(instance_name) {
return Err(format!(
"Legacy OAuth instance name contains invalid characters or exceeds max length ({LEGACY_STATE_MAX_LEN})"
));
}
tracing::debug!(
flow_id,
instance_name,
"Decoded legacy prefixed OAuth state"
);
return Ok(DecodedHostedOAuthState {
flow_id: flow_id.to_string(),
instance_name: if instance_name.is_empty() {
@@ -668,6 +715,9 @@ pub fn decode_hosted_oauth_state(state: &str) -> Result<DecodedHostedOAuthState,
return Err("Hosted OAuth state is empty".to_string());
}
validate_legacy_flow_id(state)?;
tracing::debug!(flow_id = state, "Decoded legacy raw OAuth state");
Ok(DecodedHostedOAuthState {
flow_id: state.to_string(),
instance_name: None,
@@ -1734,13 +1784,13 @@ mod tests {
fn test_decode_hosted_oauth_state_accepts_legacy_formats() {
use crate::cli::oauth_defaults::decode_hosted_oauth_state;
let decoded = decode_hosted_oauth_state("kind-deer:abc123").expect("legacy prefixed");
assert_eq!(decoded.flow_id, "abc123");
let decoded = decode_hosted_oauth_state("kind-deer:abc12345").expect("legacy prefixed");
assert_eq!(decoded.flow_id, "abc12345");
assert_eq!(decoded.instance_name.as_deref(), Some("kind-deer"));
assert!(decoded.is_legacy);
let decoded = decode_hosted_oauth_state("abc123").expect("legacy raw");
assert_eq!(decoded.flow_id, "abc123");
let decoded = decode_hosted_oauth_state("abc12345").expect("legacy raw");
assert_eq!(decoded.flow_id, "abc12345");
assert_eq!(decoded.instance_name, None);
assert!(decoded.is_legacy);
}
@@ -1864,4 +1914,72 @@ mod tests {
assert_eq!(decoded_no_instance.instance_name, None);
assert!(!decoded_no_instance.is_legacy);
}
/// Legacy flow IDs that are too short must be rejected (#1443).
#[test]
fn test_legacy_state_rejects_short_flow_id() {
use crate::cli::oauth_defaults::decode_hosted_oauth_state;
let err = decode_hosted_oauth_state("abc").expect_err("short raw flow_id");
assert!(err.contains("too short"), "unexpected error: {err}");
let err = decode_hosted_oauth_state("inst:abc").expect_err("short prefixed flow_id");
assert!(err.contains("too short"), "unexpected error: {err}");
}
/// Legacy flow IDs with invalid characters must be rejected (#1443).
#[test]
fn test_legacy_state_rejects_invalid_characters() {
use crate::cli::oauth_defaults::decode_hosted_oauth_state;
let err = decode_hosted_oauth_state("flow id with spaces!").expect_err("spaces in flow_id");
assert!(
err.contains("invalid characters"),
"unexpected error: {err}"
);
let err = decode_hosted_oauth_state("inst:flow/id?bad=yes")
.expect_err("special chars in prefixed flow_id");
assert!(
err.contains("invalid characters"),
"unexpected error: {err}"
);
}
/// Legacy instance names with invalid characters must be rejected (#1444).
#[test]
fn test_legacy_state_rejects_invalid_instance_name() {
use crate::cli::oauth_defaults::decode_hosted_oauth_state;
let err = decode_hosted_oauth_state("bad instance!:valid-flow-id-12345")
.expect_err("invalid instance name");
assert!(err.contains("instance name"), "unexpected error: {err}");
}
/// Excessively long legacy flow IDs must be rejected (#1443).
#[test]
fn test_legacy_state_rejects_oversized_flow_id() {
use crate::cli::oauth_defaults::decode_hosted_oauth_state;
let long_id = "a".repeat(200);
let err = decode_hosted_oauth_state(&long_id).expect_err("oversized flow_id");
assert!(err.contains("too long"), "unexpected error: {err}");
}
/// Valid legacy flow IDs at boundary lengths are accepted.
#[test]
fn test_legacy_state_accepts_boundary_lengths() {
use crate::cli::oauth_defaults::decode_hosted_oauth_state;
// Exactly 8 chars (minimum)
let decoded = decode_hosted_oauth_state("abcd1234").expect("8-char flow_id");
assert_eq!(decoded.flow_id, "abcd1234");
assert!(decoded.is_legacy);
// Exactly 128 chars (maximum)
let max_id = "a".repeat(128);
let decoded = decode_hosted_oauth_state(&max_id).expect("128-char flow_id");
assert_eq!(decoded.flow_id, max_id);
assert!(decoded.is_legacy);
}
}
+17 -13
View File
@@ -17,7 +17,6 @@ mod workspace;
use std::path::Path;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use async_trait::async_trait;
use chrono::{DateTime, NaiveDateTime, Utc};
@@ -34,8 +33,6 @@ use crate::workspace::MemoryDocument;
use crate::db::libsql_migrations;
static NAIVE_TIMESTAMP_LOGGED: AtomicBool = AtomicBool::new(false);
/// Explicit column list for routines table (matches positional access in `row_to_routine_libsql`).
pub(crate) const ROUTINE_COLUMNS: &str = "\
id, name, description, user_id, enabled, \
@@ -167,13 +164,11 @@ impl LibSqlBackend {
///
/// Returns an error if none of the formats match.
pub(crate) fn parse_timestamp(s: &str) -> Result<DateTime<Utc>, String> {
let log_naive_timestamp_once = || {
if !NAIVE_TIMESTAMP_LOGGED.swap(true, Ordering::Relaxed) {
tracing::debug!(
timestamp = %s,
"parsed naive timestamp without timezone; assuming UTC for backward compatibility"
);
}
let log_naive_timestamp = || {
tracing::warn!(
timestamp = %s,
"parsed naive timestamp, assuming UTC — consider migrating to RFC 3339"
);
};
// RFC 3339 (our canonical write format)
@@ -182,12 +177,12 @@ pub(crate) fn parse_timestamp(s: &str) -> Result<DateTime<Utc>, String> {
}
// Naive with fractional seconds (legacy or SQLite datetime() output)
if let Ok(ndt) = NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S%.f") {
log_naive_timestamp_once();
log_naive_timestamp();
return Ok(ndt.and_utc());
}
// Naive without fractional seconds (legacy format)
if let Ok(ndt) = NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S") {
log_naive_timestamp_once();
log_naive_timestamp();
return Ok(ndt.and_utc());
}
Err(format!("unparseable timestamp: {:?}", s))
@@ -439,7 +434,7 @@ mod tests {
use chrono::{TimeZone, Utc};
use crate::db::Database;
use crate::db::libsql::{LibSqlBackend, normalize_notify_user, parse_timestamp};
use crate::db::libsql::{LibSqlBackend, fmt_ts, normalize_notify_user, parse_timestamp};
#[test]
fn test_normalize_notify_user_treats_legacy_default_as_missing() {
@@ -468,6 +463,15 @@ mod tests {
assert_eq!(naive_without_millis, expected);
}
#[test]
fn test_fmt_ts_roundtrips_through_parse_timestamp() {
let original = Utc.with_ymd_and_hms(2026, 6, 15, 8, 30, 45).unwrap()
+ chrono::Duration::milliseconds(123);
let formatted = fmt_ts(&original);
let parsed = parse_timestamp(&formatted).unwrap();
assert_eq!(parsed, original);
}
#[tokio::test]
async fn test_libsql_now_format_is_rfc3339_and_parseable() {
let backend = LibSqlBackend::new_memory().await.unwrap();
+1 -1
View File
@@ -981,7 +981,7 @@ mod tests {
assert!(alice_stats.last_active_at.is_some());
// Bob has no LLM calls so doesn't appear in summary stats
assert!(stats.iter().find(|s| s.user_id == "bob").is_none());
assert!(!stats.iter().any(|s| s.user_id == "bob"));
// Filter to single user
let alice_only = db.user_summary_stats(Some("alice")).await.unwrap();
+649 -1
View File
@@ -53,6 +53,38 @@ struct HostedOAuthFlowStart {
flow: crate::cli::oauth_defaults::PendingOAuthFlow,
}
#[derive(Debug, Default)]
struct SecretCleanupPlan {
base_secrets: HashSet<String>,
companion_secrets: HashMap<String, HashSet<String>>,
}
impl SecretCleanupPlan {
fn add_base_secret(&mut self, secret_name: impl AsRef<str>) {
self.base_secrets
.insert(secret_name.as_ref().to_lowercase());
}
fn add_companion_secret(
&mut self,
base_secret_name: impl AsRef<str>,
companion_secret_name: impl AsRef<str>,
) {
self.companion_secrets
.entry(base_secret_name.as_ref().to_lowercase())
.or_default()
.insert(companion_secret_name.as_ref().to_lowercase());
}
}
fn oauth_refresh_secret_name(secret_name: &str) -> String {
format!("{}_refresh_token", secret_name.to_lowercase())
}
fn oauth_scopes_secret_name(secret_name: &str) -> String {
format!("{}_scopes", secret_name.to_lowercase())
}
fn normalize_oauth_callback_path(path: &str) -> String {
let trimmed_path = path.trim_end_matches('/');
if trimmed_path.is_empty() {
@@ -1602,6 +1634,10 @@ impl ExtensionManager {
match kind {
ExtensionKind::McpServer => {
let cleanup_plan = self
.collect_secret_cleanup_plan(name, kind, user_id)
.await?;
// Unregister tools with this server's prefix
let tool_names: Vec<String> = self
.tool_registry
@@ -1623,6 +1659,9 @@ impl ExtensionManager {
.await
.map_err(|e| ExtensionError::Config(e.to_string()))?;
self.cleanup_uninstalled_extension_secrets(cleanup_plan, user_id)
.await;
Ok(format!(
"Removed MCP server '{}' and {} tool(s)",
name,
@@ -1630,6 +1669,10 @@ impl ExtensionManager {
))
}
ExtensionKind::WasmTool => {
let cleanup_plan = self
.collect_secret_cleanup_plan(name, kind, user_id)
.await?;
// Unregister from tool registry
self.tool_registry.unregister(name).await;
@@ -1674,9 +1717,16 @@ impl ExtensionManager {
let _ = tokio::fs::remove_file(&cap_path).await;
}
self.cleanup_uninstalled_extension_secrets(cleanup_plan, user_id)
.await;
Ok(format!("Removed WASM tool '{}'", name))
}
ExtensionKind::WasmChannel => {
let cleanup_plan = self
.collect_secret_cleanup_plan(name, kind, user_id)
.await?;
// Remove from active set and persist
self.active_channel_names.write().await.remove(name);
self.persist_active_channels(user_id).await;
@@ -1702,6 +1752,9 @@ impl ExtensionManager {
let _ = tokio::fs::remove_file(&cap_path).await;
}
self.cleanup_uninstalled_extension_secrets(cleanup_plan, user_id)
.await;
Ok(format!(
"Removed channel '{}'. Restart IronClaw for the change to take effect.",
name
@@ -2999,6 +3052,258 @@ impl ExtensionManager {
crate::tools::wasm::CapabilitiesFile::from_bytes(&cap_bytes).ok()
}
async fn load_channel_capabilities(
&self,
name: &str,
) -> Option<crate::channels::wasm::ChannelCapabilitiesFile> {
let cap_path = self
.wasm_channels_dir
.join(format!("{}.capabilities.json", name));
let cap_bytes = tokio::fs::read(&cap_path).await.ok()?;
crate::channels::wasm::ChannelCapabilitiesFile::from_bytes(&cap_bytes).ok()
}
async fn collect_secret_cleanup_plan(
&self,
name: &str,
kind: ExtensionKind,
user_id: &str,
) -> Result<SecretCleanupPlan, ExtensionError> {
let mut plan = SecretCleanupPlan::default();
match kind {
ExtensionKind::WasmTool => {
if let Some(cap) = self.load_tool_capabilities(name).await {
for secret_name in Self::tool_secret_names(&cap) {
plan.add_base_secret(secret_name);
}
if let Some(auth) = cap.auth {
plan.add_base_secret(&auth.secret_name);
plan.add_companion_secret(
&auth.secret_name,
oauth_refresh_secret_name(&auth.secret_name),
);
plan.add_companion_secret(
&auth.secret_name,
oauth_scopes_secret_name(&auth.secret_name),
);
}
}
}
ExtensionKind::WasmChannel => {
if let Some(cap) = self.load_channel_capabilities(name).await {
for secret_name in Self::channel_secret_names(&cap) {
plan.add_base_secret(secret_name);
}
}
}
ExtensionKind::McpServer => {
let server = self
.get_mcp_server(name, user_id)
.await
.map_err(|e| ExtensionError::Config(e.to_string()))?;
let token_secret_name = server.token_secret_name();
plan.add_base_secret(&token_secret_name);
plan.add_base_secret(server.client_id_secret_name());
// MCP OAuth can persist companion secrets through two paths:
// the MCP auth helper uses `mcp_<name>_refresh_token`, while the
// hosted gateway callback stores companions alongside the access
// token secret (`<token_secret>_refresh_token` / `_scopes`).
plan.add_companion_secret(&token_secret_name, server.refresh_token_secret_name());
plan.add_companion_secret(
&token_secret_name,
oauth_refresh_secret_name(&token_secret_name),
);
plan.add_companion_secret(
&token_secret_name,
oauth_scopes_secret_name(&token_secret_name),
);
}
ExtensionKind::ChannelRelay => {}
}
Ok(plan)
}
async fn cleanup_uninstalled_extension_secrets(&self, plan: SecretCleanupPlan, user_id: &str) {
let referenced_secrets = match self.collect_referenced_secret_names(user_id).await {
Ok(secret_names) => secret_names,
Err(error) => {
tracing::warn!(
user_id,
error,
"Failed to determine which secrets are still referenced; keeping secrets"
);
return;
}
};
for base_secret in &plan.base_secrets {
if referenced_secrets.contains(base_secret) {
continue;
}
self.delete_secret_best_effort(user_id, base_secret).await;
if let Some(companion_secrets) = plan.companion_secrets.get(base_secret) {
for companion_secret in companion_secrets {
if !referenced_secrets.contains(companion_secret) {
self.delete_secret_best_effort(user_id, companion_secret)
.await;
}
}
}
}
}
async fn delete_secret_best_effort(&self, user_id: &str, secret_name: &str) {
if let Err(error) = self.secrets.delete(user_id, secret_name).await {
tracing::warn!(
user_id,
secret_name,
error = %error,
"Failed to delete secret while uninstalling extension"
);
}
}
async fn collect_referenced_secret_names(
&self,
user_id: &str,
) -> Result<HashSet<String>, String> {
let mut referenced_secret_names = HashSet::new();
let tools = discover_tools(&self.wasm_tools_dir)
.await
.map_err(|e| format!("discover tools: {e}"))?;
for (tool_name, discovered_tool) in &tools {
let cap = self
.load_tool_capabilities(tool_name)
.await
.ok_or_else(|| {
let path = discovered_tool
.capabilities_path
.as_ref()
.map(|path| path.display().to_string())
.unwrap_or_else(|| format!("{} (missing)", tool_name));
format!("load tool capabilities for {tool_name}: {path}")
})?;
referenced_secret_names.extend(Self::tool_secret_names(&cap));
}
let channels = crate::channels::wasm::discover_channels(&self.wasm_channels_dir)
.await
.map_err(|e| format!("discover channels: {e}"))?;
for (channel_name, discovered_channel) in &channels {
let cap = self
.load_channel_capabilities(channel_name)
.await
.ok_or_else(|| {
let path = discovered_channel
.capabilities_path
.as_ref()
.map(|path| path.display().to_string())
.unwrap_or_else(|| format!("{} (missing)", channel_name));
format!("load channel capabilities for {channel_name}: {path}")
})?;
referenced_secret_names.extend(Self::channel_secret_names(&cap));
}
let mcp_servers = self
.load_mcp_servers(user_id)
.await
.map_err(|e| format!("load MCP servers: {e}"))?;
for server in &mcp_servers.servers {
referenced_secret_names.extend(Self::mcp_server_secret_names(server));
}
Ok(referenced_secret_names)
}
fn tool_secret_names(cap: &crate::tools::wasm::CapabilitiesFile) -> HashSet<String> {
let mut names = HashSet::new();
if let Some(auth) = &cap.auth {
names.insert(auth.secret_name.to_lowercase());
}
if let Some(setup) = &cap.setup {
names.extend(
setup
.required_secrets
.iter()
.map(|secret| secret.name.to_lowercase()),
);
}
if let Some(http) = &cap.http {
names.extend(
http.credentials
.values()
.map(|credential| credential.secret_name.to_lowercase()),
);
}
if let Some(webhook) = &cap.webhook {
if let Some(secret_name) = &webhook.secret_name {
names.insert(secret_name.to_lowercase());
}
if let Some(secret_name) = &webhook.signature_key_secret_name {
names.insert(secret_name.to_lowercase());
}
if let Some(secret_name) = &webhook.hmac_secret_name {
names.insert(secret_name.to_lowercase());
}
}
names
}
fn channel_secret_names(
cap: &crate::channels::wasm::ChannelCapabilitiesFile,
) -> HashSet<String> {
let mut names: HashSet<String> = cap
.setup
.required_secrets
.iter()
.map(|secret| secret.name.to_lowercase())
.collect();
if let Some(http) = cap.capabilities.tool.http.as_ref() {
names.extend(
http.credentials
.values()
.map(|credential| credential.secret_name.to_lowercase()),
);
}
if let Some(webhook) = cap
.capabilities
.channel
.as_ref()
.and_then(|channel| channel.webhook.as_ref())
{
if webhook.secret_header.is_some() || webhook.secret_name.is_some() {
names.insert(cap.webhook_secret_name().to_lowercase());
}
if let Some(secret_name) = cap.signature_key_secret_name() {
names.insert(secret_name.to_lowercase());
}
if let Some(secret_name) = cap.hmac_secret_name() {
names.insert(secret_name.to_lowercase());
}
}
names
}
fn mcp_server_secret_names(server: &McpServerConfig) -> HashSet<String> {
[
server.token_secret_name().to_lowercase(),
server.client_id_secret_name().to_lowercase(),
]
.into_iter()
.collect()
}
/// Collect merged OAuth scopes from all installed tools sharing the same secret_name.
///
/// When multiple tools share an OAuth provider (e.g., google-calendar and google-drive
@@ -6033,6 +6338,8 @@ mod tests {
ExtensionError, ExtensionKind, ExtensionSource, InstallResult, VerificationChallenge,
};
use crate::pairing::PairingStore;
use crate::secrets::CreateSecretParams;
use crate::tools::mcp::McpServerConfig;
fn require(condition: bool, message: impl Into<String>) -> Result<(), String> {
if condition {
@@ -6353,6 +6660,38 @@ mod tests {
tools_dir
}
fn write_test_channel(
dir: &std::path::Path,
name: &str,
capabilities_json: &str,
) -> std::path::PathBuf {
let channels_dir = dir.join("channels");
std::fs::create_dir_all(&channels_dir).expect("channels dir");
std::fs::write(
channels_dir.join(format!("{name}.wasm")),
b"not-a-real-wasm",
)
.expect("wasm");
std::fs::write(
channels_dir.join(format!("{name}.capabilities.json")),
capabilities_json,
)
.expect("capabilities");
channels_dir
}
async fn store_test_secret(
manager: &crate::extensions::manager::ExtensionManager,
name: &str,
value: &str,
) {
manager
.secrets
.create("test", CreateSecretParams::new(name, value))
.await
.expect("store secret");
}
#[test]
fn test_setting_value_is_present() {
assert!(
@@ -7423,7 +7762,13 @@ mod tests {
// Regression: remove() only checked channel_runtime for shutdown, missing
// relay-only mode where only relay_channel_manager is set.
let dir = tempfile::tempdir().expect("temp dir");
let mgr = make_test_manager(None, dir.path().to_path_buf());
let (store, _db_dir) = make_test_store().await;
let mgr = make_test_manager_with_dirs(
None,
dir.path().join("tools"),
dir.path().join("channels"),
Some(store),
);
// Set up relay channel manager with a stub channel
let cm = Arc::new(crate::channels::ChannelManager::new());
@@ -7450,6 +7795,8 @@ mod tests {
.await
.expect("store team_id");
}
store_test_secret(&mgr, "relay:slack-relay:oauth_state", "nonce").await;
store_test_secret(&mgr, "relay:slack-relay:stream_token", "legacy-token").await;
// Verify channel exists before removal
assert!(cm.get_channel("slack-relay").await.is_some());
@@ -7478,6 +7825,30 @@ mod tests {
cm.get_channel("slack-relay").await.is_none(),
"relay channel should be removed from the channel manager"
);
assert!(
!mgr.secrets
.exists("test", "relay:slack-relay:oauth_state")
.await
.expect("oauth state exists query"),
"relay oauth_state secret should be removed"
);
assert!(
!mgr.secrets
.exists("test", "relay:slack-relay:stream_token")
.await
.expect("stream token exists query"),
"relay legacy stream token should be removed"
);
assert_eq!(
mgr.store
.as_ref()
.expect("store")
.get_setting("test", "relay:slack-relay:team_id")
.await
.expect("team_id query"),
None,
"relay team_id setting should be removed"
);
}
#[tokio::test]
@@ -7585,6 +7956,185 @@ mod tests {
);
}
#[tokio::test]
async fn test_remove_wasm_tool_deletes_unique_secrets() {
let dir = tempfile::tempdir().expect("temp dir");
let tools_dir = write_test_tool(
dir.path(),
"github",
r#"{
"name": "github",
"auth": { "secret_name": "github_token" },
"setup": {
"required_secrets": [
{ "name": "github_client_secret", "prompt": "GitHub client secret for testing cleanup behavior." }
]
},
"http": {
"credentials": {
"service_token": {
"secret_name": "github_service_token",
"location": { "type": "bearer" }
}
}
},
"webhook": {
"hmac_secret_name": "github_webhook_secret"
}
}"#,
);
let mgr = make_test_manager_with_dirs(None, tools_dir, dir.path().join("channels"), None);
store_test_secret(&mgr, "github_token", "access-token").await;
store_test_secret(&mgr, "github_token_refresh_token", "refresh-token").await;
store_test_secret(&mgr, "github_token_scopes", "repo workflow").await;
store_test_secret(&mgr, "github_client_secret", "client-secret").await;
store_test_secret(&mgr, "github_service_token", "service-token").await;
store_test_secret(&mgr, "github_webhook_secret", "webhook-secret").await;
mgr.remove("github", "test")
.await
.expect("remove should succeed");
for secret_name in [
"github_token",
"github_token_refresh_token",
"github_token_scopes",
"github_client_secret",
"github_service_token",
"github_webhook_secret",
] {
assert!(
!mgr.secrets
.exists("test", secret_name)
.await
.expect("exists query"),
"secret {secret_name} should be deleted"
);
}
}
#[tokio::test]
async fn test_remove_wasm_tool_keeps_secrets_when_other_tool_capabilities_missing() {
let dir = tempfile::tempdir().expect("temp dir");
let tools_dir = write_test_tool(
dir.path(),
"github",
r#"{
"name": "github",
"auth": { "secret_name": "shared_token" }
}"#,
);
std::fs::write(tools_dir.join("broken.wasm"), b"fake-tool").expect("write tool");
let mgr = make_test_manager_with_dirs(None, tools_dir, dir.path().join("channels"), None);
store_test_secret(&mgr, "shared_token", "access-token").await;
store_test_secret(&mgr, "shared_token_refresh_token", "refresh-token").await;
store_test_secret(&mgr, "shared_token_scopes", "repo").await;
mgr.remove("github", "test")
.await
.expect("remove should succeed");
for secret_name in [
"shared_token",
"shared_token_refresh_token",
"shared_token_scopes",
] {
assert!(
mgr.secrets
.exists("test", secret_name)
.await
.expect("exists query"),
"secret {secret_name} should be retained when reference detection is uncertain"
);
}
}
#[tokio::test]
async fn test_remove_wasm_tool_keeps_shared_secrets_until_last_extension() {
let dir = tempfile::tempdir().expect("temp dir");
write_test_tool(
dir.path(),
"google-calendar",
r#"{
"name": "google-calendar",
"auth": { "secret_name": "google_oauth_token" },
"setup": {
"required_secrets": [
{ "name": "google_oauth_client_id", "prompt": "Google OAuth client id for cleanup testing." },
{ "name": "google_oauth_client_secret", "prompt": "Google OAuth client secret for cleanup testing." }
]
}
}"#,
);
let tools_dir = write_test_tool(
dir.path(),
"google-drive",
r#"{
"name": "google-drive",
"auth": { "secret_name": "google_oauth_token" },
"setup": {
"required_secrets": [
{ "name": "google_oauth_client_id", "prompt": "Google OAuth client id for cleanup testing." },
{ "name": "google_oauth_client_secret", "prompt": "Google OAuth client secret for cleanup testing." }
]
}
}"#,
);
let mgr = make_test_manager_with_dirs(None, tools_dir, dir.path().join("channels"), None);
for (secret_name, value) in [
("google_oauth_token", "access-token"),
("google_oauth_token_refresh_token", "refresh-token"),
("google_oauth_token_scopes", "calendar drive"),
("google_oauth_client_id", "client-id"),
("google_oauth_client_secret", "client-secret"),
] {
store_test_secret(&mgr, secret_name, value).await;
}
mgr.remove("google-calendar", "test")
.await
.expect("first remove should succeed");
for secret_name in [
"google_oauth_token",
"google_oauth_token_refresh_token",
"google_oauth_token_scopes",
"google_oauth_client_id",
"google_oauth_client_secret",
] {
assert!(
mgr.secrets
.exists("test", secret_name)
.await
.expect("exists query"),
"shared secret {secret_name} should remain while google-drive is still installed"
);
}
mgr.remove("google-drive", "test")
.await
.expect("second remove should succeed");
for secret_name in [
"google_oauth_token",
"google_oauth_token_refresh_token",
"google_oauth_token_scopes",
"google_oauth_client_id",
"google_oauth_client_secret",
] {
assert!(
!mgr.secrets
.exists("test", secret_name)
.await
.expect("exists query"),
"shared secret {secret_name} should be deleted after the last tool is removed"
);
}
}
#[tokio::test]
async fn test_remove_wasm_channel_clears_activation_error_and_deletes_files() {
let dir = tempfile::tempdir().expect("temp dir");
@@ -7619,6 +8169,104 @@ mod tests {
);
}
#[tokio::test]
async fn test_remove_wasm_channel_deletes_setup_secrets() {
let dir = tempfile::tempdir().expect("temp dir");
let channels_dir = write_test_channel(
dir.path(),
"telegram",
r#"{
"type": "channel",
"name": "telegram",
"setup": {
"required_secrets": [
{
"name": "telegram_bot_token",
"prompt": "Telegram bot token used to verify uninstall cleanup behavior."
}
]
},
"capabilities": {
"http": {
"credentials": {
"tenant_token": {
"secret_name": "telegram_service_token",
"location": { "type": "bearer" }
}
}
},
"channel": {
"webhook": {
"secret_header": "X-Telegram-Bot-Api-Secret-Token",
"secret_name": "telegram_webhook_secret"
}
}
}
}"#,
);
let mgr = make_test_manager_with_dirs(None, dir.path().join("tools"), channels_dir, None);
store_test_secret(&mgr, "telegram_bot_token", "123:telegram-token").await;
store_test_secret(&mgr, "telegram_service_token", "tenant-service-token").await;
store_test_secret(&mgr, "telegram_webhook_secret", "webhook-secret").await;
mgr.remove("telegram", "test")
.await
.expect("remove should succeed");
for secret_name in [
"telegram_bot_token",
"telegram_service_token",
"telegram_webhook_secret",
] {
assert!(
!mgr.secrets
.exists("test", secret_name)
.await
.expect("exists query"),
"channel secret {secret_name} should be deleted"
);
}
}
#[tokio::test]
async fn test_remove_mcp_server_deletes_stored_secrets() {
let dir = tempfile::tempdir().expect("temp dir");
let (store, _db_dir) = make_test_store().await;
let mgr = make_test_manager_with_dirs(
None,
dir.path().join("tools"),
dir.path().join("channels"),
Some(Arc::clone(&store)),
);
let server = McpServerConfig::new("notion", "https://example.com/mcp");
mgr.add_mcp_server(server.clone(), "test")
.await
.expect("add mcp server");
store_test_secret(&mgr, &server.token_secret_name(), "access-token").await;
store_test_secret(&mgr, &server.refresh_token_secret_name(), "refresh-token").await;
store_test_secret(&mgr, &server.client_id_secret_name(), "client-id").await;
mgr.remove("notion", "test")
.await
.expect("remove should succeed");
for secret_name in [
server.token_secret_name(),
server.refresh_token_secret_name(),
server.client_id_secret_name(),
] {
assert!(
!mgr.secrets
.exists("test", &secret_name)
.await
.expect("exists query"),
"MCP secret {secret_name} should be deleted"
);
}
}
#[test]
fn test_sanitize_url_with_query_params() {
let url = "https://api.example.com/path?api_key=secret123&token=abc";
+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(_)
+4 -2
View File
@@ -124,8 +124,10 @@ impl WasmToolLoader {
let wasm_bytes = fs::read(wasm_path).await?;
// Read capabilities (optional) and extract OAuth refresh config
// and tool description. Parameter schema is auto-derived from the
// WASM module's schema() export (see WasmToolSchemas::compact_schema).
// and tool description. Parameter schema is NOT read from the
// capabilities file — it is auto-derived from the WASM module's
// schema() export at prepare time (see WasmToolSchemas::compact_schema),
// so no schema override is needed here.
let (capabilities, oauth_refresh, description) = if let Some(cap_path) = capabilities_path {
if cap_path.exists() {
let cap_bytes = fs::read(cap_path).await?;
+92 -5
View File
@@ -759,14 +759,31 @@ impl WasmToolSchemas {
}
let kept: serde_json::Map<String, serde_json::Value> = all_properties
.into_iter()
.iter()
.filter(|(name, prop)| {
required.contains(name) || prop.get("enum").is_some() || prop.get("const").is_some()
required.contains(name.as_str())
|| prop.get("enum").is_some()
|| prop.get("const").is_some()
})
.map(|(k, v)| (k.clone(), v.clone()))
.collect();
if kept.is_empty() {
return Self::permissive_schema();
// When the schema has typed properties but none survived the
// required/enum filter, include all typed properties so the LLM
// sees meaningful parameter hints instead of permissive `{}`.
let typed: serde_json::Map<String, serde_json::Value> = all_properties
.into_iter()
.filter(|(_, prop)| schema_is_typed_property(prop))
.collect();
if typed.is_empty() {
return Self::permissive_schema();
}
return serde_json::json!({
"type": "object",
"properties": typed,
"additionalProperties": true,
});
}
let kept_required: Vec<serde_json::Value> = required
@@ -1991,6 +2008,58 @@ mod tests {
);
}
#[tokio::test]
async fn test_typed_schema_without_required_is_advertised() {
// Regression test for #1303: when a WASM tool exports a typed schema
// with no required/enum fields, the advertised schema should still
// contain the typed properties instead of falling back to permissive {}.
let discovery_schema = serde_json::json!({
"type": "object",
"properties": {
"query": { "type": "string" },
"limit": { "type": "integer" }
}
});
let runtime = Arc::new(WasmToolRuntime::new(WasmRuntimeConfig::for_testing()).unwrap());
let prepared = runtime
.prepare("typed_search", b"\0asm\x0d\0\x01\0", None)
.await
.unwrap();
let mut wrapper =
super::WasmToolWrapper::new(Arc::clone(&runtime), prepared, Capabilities::default());
wrapper.schemas = super::WasmToolSchemas::new(discovery_schema.clone());
wrapper.description = "Typed search tool".to_string();
let advertised = wrapper.parameters_schema();
let props = advertised["properties"].as_object().unwrap();
// Both typed properties should be preserved in the advertised schema
assert!(
props.contains_key("query"),
"advertised schema should contain 'query' property"
);
assert!(
props.contains_key("limit"),
"advertised schema should contain 'limit' property"
);
assert_eq!(props.len(), 2);
// The schema should NOT be permissive
assert!(
!super::WasmToolSchemas::is_permissive_schema(&advertised),
"advertised schema should not be permissive when typed properties exist"
);
// No tool_info hint needed since typed properties are visible
let schema = wrapper.schema();
assert!(
!schema.description.contains("tool_info"),
"description should not contain tool_info hint: {}",
schema.description
);
}
#[test]
fn test_compact_schema_keeps_required_and_enum_properties() {
let schema = serde_json::json!({
@@ -2028,8 +2097,8 @@ mod tests {
}
#[test]
fn test_compact_schema_falls_back_to_permissive_when_empty() {
// No required, no enum → permissive fallback
fn test_compact_schema_preserves_typed_properties_when_no_required() {
// No required, no enum, but typed properties → keep all typed props
let schema = serde_json::json!({
"type": "object",
"properties": {
@@ -2038,6 +2107,24 @@ mod tests {
}
});
let compacted = super::WasmToolSchemas::compact_schema(&schema);
let props = compacted["properties"].as_object().unwrap();
assert_eq!(props.len(), 2);
assert!(props.contains_key("query"));
assert!(props.contains_key("limit"));
assert_eq!(compacted["additionalProperties"], true);
}
#[test]
fn test_compact_schema_falls_back_to_permissive_when_no_typed_properties() {
// Properties with no type info → permissive fallback
let schema = serde_json::json!({
"type": "object",
"properties": {
"data": {}
}
});
let compacted = super::WasmToolSchemas::compact_schema(&schema);
assert!(compacted["properties"].as_object().unwrap().is_empty());
}
+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(),
}
));
}
}
+2
View File
@@ -53,6 +53,7 @@ HEADED=1 pytest scenarios/
| `test_skills.py` | Skills tab UI visibility, ClawHub search (skipped if registry unreachable), install + remove lifecycle |
| `test_sse_reconnect.py` | SSE reconnects after programmatic `eventSource.close()` + `connectSSE()`; history is reloaded after reconnect |
| `test_tool_approval.py` | Approval card appears, buttons disable on approve/deny, parameters toggle via `page.evaluate("showApproval(...)")`; the waiting-approval regression uses a real HTTP tool call |
| `test_extension_uninstall_cleanup.py` | Real install/setup/remove coverage for WASM tools, WASM channels, OAuth-backed shared Google tools, and MCP servers; verifies uninstall deletes stored secrets from the libSQL `secrets` table while preserving shared credentials until the last referencing extension is removed |
| `test_oauth_refresh.py` | Hosted Gmail OAuth regression: complete setup via `/oauth/callback`, expire the stored access token in libSQL, trigger a real `gmail` tool call through `/api/chat/send`, and verify refresh goes through the mock `/oauth/refresh` proxy without forwarding `client_secret` |
## `helpers.py`
@@ -77,6 +78,7 @@ All fixtures are defined in `tests/e2e/conftest.py`. Running `pytest scenarios/`
| `mock_llm_server` | Starts `mock_llm.py --port 0`, reads the assigned port from stdout, waits for `/v1/models` to return 200. Yields the base URL. |
| `ironclaw_server` | Starts the ironclaw binary with a minimal env (see below), waits for `/api/health` (timeout 60s). Yields the base URL. On teardown sends **SIGINT** (not SIGTERM) so the tokio ctrl_c handler triggers a graceful shutdown and LLVM coverage data is flushed. |
| `hosted_oauth_refresh_server` | Starts a second ironclaw instance with a dedicated libSQL DB and `GOOGLE_OAUTH_CLIENT_ID=hosted-google-client-id`, while still pointing `IRONCLAW_OAUTH_EXCHANGE_URL` at `mock_llm.py`. Yields a dict with `base_url`, `db_path`, `gateway_user_id`, and `mock_llm_url` for the hosted refresh regression scenario. |
| `extension_cleanup_server` | Starts an isolated ironclaw instance with its own temp DB/home/WASM dirs, `SECRETS_MASTER_KEY`, and hosted-style OAuth env so uninstall-cleanup scenarios can inspect the `secrets` table without interfering with the shared E2E server state. |
| `browser` | Launches a single Chromium instance (headless by default; set `HEADED=1` for headed). Shared across all tests. |
### Function-scoped fixtures
+109
View File
@@ -443,6 +443,115 @@ async def hosted_oauth_refresh_server(
home_tmpdir.cleanup()
@pytest.fixture(scope="session")
async def extension_cleanup_server(
ironclaw_binary,
mock_llm_server,
):
"""Start an isolated ironclaw instance for uninstall secret cleanup E2E tests."""
reserved = _reserve_loopback_sockets(2)
db_tmpdir = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-cleanup-db-")
home_tmpdir = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-cleanup-home-")
tools_tmpdir = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-cleanup-tools-")
channels_tmpdir = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-cleanup-channels-")
try:
gateway_port = reserved[0].getsockname()[1]
http_port = reserved[1].getsockname()[1]
for sock in reserved:
if sock.fileno() != -1:
sock.close()
db_path = os.path.join(db_tmpdir.name, "extension-cleanup.db")
home_dir = home_tmpdir.name
env = {
"PATH": os.environ.get("PATH", "/usr/bin:/bin"),
"HOME": home_dir,
"IRONCLAW_BASE_DIR": os.path.join(home_dir, ".ironclaw"),
"RUST_LOG": "ironclaw=info",
"RUST_BACKTRACE": "1",
"IRONCLAW_OWNER_ID": OWNER_SCOPE_ID,
"GATEWAY_ENABLED": "true",
"GATEWAY_HOST": "127.0.0.1",
"GATEWAY_PORT": str(gateway_port),
"GATEWAY_AUTH_TOKEN": AUTH_TOKEN,
"GATEWAY_USER_ID": OWNER_SCOPE_ID,
"HTTP_HOST": "127.0.0.1",
"HTTP_PORT": str(http_port),
"HTTP_WEBHOOK_SECRET": HTTP_WEBHOOK_SECRET,
"CLI_ENABLED": "false",
"LLM_BACKEND": "openai_compatible",
"LLM_BASE_URL": mock_llm_server,
"LLM_MODEL": "mock-model",
"DATABASE_BACKEND": "libsql",
"LIBSQL_PATH": db_path,
"SECRETS_MASTER_KEY": "0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef",
"SANDBOX_ENABLED": "false",
"SKILLS_ENABLED": "true",
"ROUTINES_ENABLED": "true",
"HEARTBEAT_ENABLED": "false",
"EMBEDDING_ENABLED": "false",
"WASM_ENABLED": "true",
"WASM_TOOLS_DIR": tools_tmpdir.name,
"WASM_CHANNELS_DIR": channels_tmpdir.name,
"ONBOARD_COMPLETED": "true",
"IRONCLAW_OAUTH_CALLBACK_URL": "https://oauth.test.example/oauth/callback",
"IRONCLAW_OAUTH_EXCHANGE_URL": mock_llm_server,
"GOOGLE_OAUTH_CLIENT_ID": "hosted-google-client-id",
}
_forward_coverage_env(env)
proc = await asyncio.create_subprocess_exec(
ironclaw_binary, "--no-onboard",
stdin=asyncio.subprocess.DEVNULL,
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
env=env,
)
startup_kill_attempted = False
base_url = f"http://127.0.0.1:{gateway_port}"
try:
await wait_for_ready(f"{base_url}/api/health", timeout=60)
yield {
"base_url": base_url,
"db_path": db_path,
"gateway_user_id": OWNER_SCOPE_ID,
"mock_llm_url": mock_llm_server,
}
except TimeoutError:
if proc.returncode is None:
startup_kill_attempted = True
await _stop_process(proc, timeout=2)
returncode = proc.returncode
stderr_bytes = b""
if proc.stderr:
try:
stderr_bytes = await asyncio.wait_for(proc.stderr.read(8192), timeout=2)
except asyncio.TimeoutError:
pass
stderr_text = stderr_bytes.decode("utf-8", errors="replace")
pytest.fail(
f"extension cleanup server failed to start on port {gateway_port} "
f"(returncode={returncode}).\nstderr:\n{stderr_text}"
)
finally:
if proc.returncode is None:
if startup_kill_attempted:
await _stop_process(proc, timeout=2)
else:
await _stop_process(proc, sig=signal.SIGINT, timeout=10)
if proc.returncode is None:
await _stop_process(proc, timeout=2)
finally:
for sock in reserved:
if sock.fileno() != -1:
sock.close()
db_tmpdir.cleanup()
home_tmpdir.cleanup()
tools_tmpdir.cleanup()
channels_tmpdir.cleanup()
@pytest.fixture(scope="session")
async def http_channel_server(ironclaw_server, server_ports):
"""HTTP webhook channel base URL."""
@@ -0,0 +1,266 @@
"""Extension uninstall secret cleanup E2E tests.
Exercises real install/setup/auth/remove flows and verifies the backing
secrets table is cleaned up when extensions are uninstalled.
"""
import sqlite3
from urllib.parse import parse_qs, urlparse
import httpx
from helpers import api_get, api_post
def _extract_state(auth_url: str) -> str:
parsed = urlparse(auth_url)
state = parse_qs(parsed.query).get("state", [None])[0]
assert state, f"auth_url should include state: {auth_url}"
return state
def _secret_exists(db_path: str, user_id: str, name: str) -> bool:
with sqlite3.connect(db_path) as conn:
row = conn.execute(
"SELECT 1 FROM secrets WHERE user_id = ?1 AND name = ?2 LIMIT 1",
(user_id, name),
).fetchone()
return row is not None
def _secret_names(db_path: str, user_id: str) -> set[str]:
with sqlite3.connect(db_path) as conn:
rows = conn.execute(
"SELECT name FROM secrets WHERE user_id = ?1",
(user_id,),
).fetchall()
return {row[0] for row in rows}
async def _get_extension(base_url: str, name: str) -> dict | None:
response = await api_get(base_url, "/api/extensions", timeout=15)
response.raise_for_status()
for extension in response.json().get("extensions", []):
if extension["name"] == name:
return extension
return None
async def _ensure_removed(base_url: str, name: str) -> None:
extension = await _get_extension(base_url, name)
if extension is not None:
response = await api_post(base_url, f"/api/extensions/{name}/remove", timeout=30)
assert response.status_code == 200, response.text
assert response.json().get("success") is True, response.text
async def _install_extension(
base_url: str,
name: str,
*,
kind: str | None = None,
url: str | None = None,
) -> None:
payload = {"name": name}
if kind is not None:
payload["kind"] = kind
if url is not None:
payload["url"] = url
response = await api_post(
base_url,
"/api/extensions/install",
json=payload,
timeout=180,
)
assert response.status_code == 200, response.text
assert response.json().get("success") is True, response.text
async def test_remove_wasm_tool_deletes_unique_secret(extension_cleanup_server):
server = extension_cleanup_server["base_url"]
db_path = extension_cleanup_server["db_path"]
user_id = extension_cleanup_server["gateway_user_id"]
await _ensure_removed(server, "web-search")
await _install_extension(server, "web-search")
setup_response = await api_post(
server,
"/api/extensions/web-search/setup",
json={"secrets": {"brave_api_key": "cleanup-test-key"}},
timeout=30,
)
assert setup_response.status_code == 200, setup_response.text
assert setup_response.json().get("success") is True, setup_response.text
assert _secret_exists(db_path, user_id, "brave_api_key")
remove_response = await api_post(
server,
"/api/extensions/web-search/remove",
timeout=30,
)
assert remove_response.status_code == 200, remove_response.text
assert remove_response.json().get("success") is True, remove_response.text
assert not _secret_exists(db_path, user_id, "brave_api_key")
async def test_remove_wasm_channel_deletes_setup_secrets(extension_cleanup_server):
server = extension_cleanup_server["base_url"]
db_path = extension_cleanup_server["db_path"]
user_id = extension_cleanup_server["gateway_user_id"]
await _ensure_removed(server, "discord")
await _install_extension(server, "discord", kind="wasm_channel")
setup_response = await api_post(
server,
"/api/extensions/discord/setup",
json={
"secrets": {
"discord_bot_token": "cleanup-discord-bot-token",
"discord_public_key": "cleanup-discord-public-key",
}
},
timeout=30,
)
assert setup_response.status_code == 200, setup_response.text
assert setup_response.json().get("success") is True, setup_response.text
assert _secret_exists(db_path, user_id, "discord_bot_token")
assert _secret_exists(db_path, user_id, "discord_public_key")
remove_response = await api_post(
server,
"/api/extensions/discord/remove",
timeout=30,
)
assert remove_response.status_code == 200, remove_response.text
assert remove_response.json().get("success") is True, remove_response.text
assert not _secret_exists(db_path, user_id, "discord_bot_token")
assert not _secret_exists(db_path, user_id, "discord_public_key")
async def test_remove_shared_google_oauth_secrets_after_last_tool(extension_cleanup_server):
server = extension_cleanup_server["base_url"]
db_path = extension_cleanup_server["db_path"]
user_id = extension_cleanup_server["gateway_user_id"]
await _ensure_removed(server, "gmail")
await _ensure_removed(server, "google-drive")
await _install_extension(server, "gmail")
await _install_extension(server, "google-drive")
setup_response = await api_post(
server,
"/api/extensions/gmail/setup",
json={"secrets": {}},
timeout=30,
)
assert setup_response.status_code == 200, setup_response.text
auth_url = setup_response.json().get("auth_url")
assert auth_url, setup_response.text
async with httpx.AsyncClient() as client:
callback_response = await client.get(
f"{server}/oauth/callback",
params={"code": "mock_auth_code", "state": _extract_state(auth_url)},
timeout=30,
follow_redirects=True,
)
assert callback_response.status_code == 200, callback_response.text[:400]
shared_secrets = [
"google_oauth_token",
"google_oauth_token_refresh_token",
"google_oauth_token_scopes",
]
for secret_name in shared_secrets:
assert _secret_exists(db_path, user_id, secret_name), f"expected {secret_name} to exist"
gmail_remove_response = await api_post(
server,
"/api/extensions/gmail/remove",
timeout=30,
)
assert gmail_remove_response.status_code == 200, gmail_remove_response.text
assert gmail_remove_response.json().get("success") is True, gmail_remove_response.text
for secret_name in shared_secrets:
assert _secret_exists(db_path, user_id, secret_name), (
f"{secret_name} should remain while google-drive is still installed"
)
drive_remove_response = await api_post(
server,
"/api/extensions/google-drive/remove",
timeout=30,
)
assert drive_remove_response.status_code == 200, drive_remove_response.text
assert drive_remove_response.json().get("success") is True, drive_remove_response.text
for secret_name in shared_secrets:
assert not _secret_exists(db_path, user_id, secret_name), (
f"{secret_name} should be deleted after the last Google tool is removed"
)
async def test_remove_mcp_server_deletes_stored_secrets(extension_cleanup_server):
server = extension_cleanup_server["base_url"]
db_path = extension_cleanup_server["db_path"]
user_id = extension_cleanup_server["gateway_user_id"]
mcp_url = f"{extension_cleanup_server['mock_llm_url']}/mcp"
await _ensure_removed(server, "mock-mcp")
await _install_extension(server, "mock-mcp", kind="mcp_server", url=mcp_url)
setup_response = await api_post(
server,
"/api/extensions/mock-mcp/setup",
json={"secrets": {}},
timeout=30,
)
assert setup_response.status_code == 200, setup_response.text
auth_url = setup_response.json().get("auth_url")
if auth_url is None:
activate_response = await api_post(
server,
"/api/extensions/mock-mcp/activate",
timeout=30,
)
assert activate_response.status_code == 200, activate_response.text
auth_url = activate_response.json().get("auth_url")
assert auth_url, "mock-mcp should require OAuth in E2E"
async with httpx.AsyncClient() as client:
callback_response = await client.get(
f"{server}/oauth/callback",
params={"code": "mock_mcp_code", "state": _extract_state(auth_url)},
timeout=30,
follow_redirects=True,
)
assert callback_response.status_code == 200, callback_response.text[:400]
expected_mcp_secrets = [
"mcp_mock-mcp_access_token",
"mcp_mock-mcp_client_id",
]
stored_secret_names = _secret_names(db_path, user_id)
for secret_name in expected_mcp_secrets:
assert secret_name in stored_secret_names, (
f"expected {secret_name} to exist; stored secrets were {sorted(stored_secret_names)}"
)
remove_response = await api_post(
server,
"/api/extensions/mock-mcp/remove",
timeout=30,
)
assert remove_response.status_code == 200, remove_response.text
assert remove_response.json().get("success") is True, remove_response.text
remaining_secret_names = _secret_names(db_path, user_id)
assert not any(name.startswith("mcp_mock-mcp_") for name in remaining_secret_names), (
f"mock-mcp secrets should be deleted on remove; remaining secrets were "
f"{sorted(remaining_secret_names)}"
)