From 45ec691f4cff536c2c5dc647603990532174aa97 Mon Sep 17 00:00:00 2001 From: Zaki Manian Date: Sat, 7 Mar 2026 00:30:47 -0800 Subject: [PATCH] Improve test infrastructure: StubChannel, gateway helpers, security tests, search edge cases (#623) * feat(testing): add StubChannel test double for Channel trait Adds StubChannel to src/testing.rs alongside StubLlm. Supports message injection via mpsc sender, response/status capture, and configurable health check toggling. Includes handle methods for use after ownership transfer to ChannelManager. Co-Authored-By: Claude Opus 4.6 * feat(testing): wire StubChannel into TestHarnessBuilder Add with_stub_channel() builder method that creates a StubChannel pre-registered in a ChannelManager. Tests can inject messages via the sender and verify routing through the manager. The channel field on TestHarness is Optional, defaulting to None for backward compat. Co-Authored-By: Claude Opus 4.6 * test: gate external-service tests behind integration feature flag Replace silent try_connect() skip pattern with explicit feature gating. cargo test now runs only self-contained tests. cargo test --features integration runs tests requiring PostgreSQL. Co-Authored-By: Claude Opus 4.6 * test(channels): add ChannelManager unit tests using StubChannel Cover add/start_all stream merging, respond routing, unknown channel errors, health_check_all with mixed health, empty-channels error path, and injection channel merging -- all via StubChannel test double. Co-Authored-By: Claude Opus 4.6 * docs: document test tier separation (unit/integration/live) Co-Authored-By: Claude Opus 4.6 * ci: add architecture boundary check script Grep-based checks for three architecture boundaries: - Direct database driver usage (tokio_postgres/libsql) outside src/db/ - .unwrap()/.expect() in production code (warning only) - Direct std::env::var reads outside config layer (warning only) The DB driver check is a hard violation; the other two are warnings for gradual cleanup. Run with: bash scripts/check-boundaries.sh Co-Authored-By: Claude Opus 4.6 * test(search): add RRF edge case tests for empty inputs, limits, and config modes Co-Authored-By: Claude Opus 4.6 * test(security): add regression tests for skill installer ZIP and SSRF protections Add 11 regression tests covering the security controls in skill_tools: ZIP extraction safety: - Valid SKILL.md extraction works correctly - Non-SKILL.md entries are ignored (returns error) - Path traversal entries (../../SKILL.md) do not match - Nested path entries (subdir/SKILL.md) do not match - Oversized entries (>1MB uncompressed) are rejected SSRF prevention: - Loopback addresses (127.0.0.1) are blocked - Private ranges (10.x, 172.16.x, 192.168.x) are blocked - Link-local addresses (169.254.x) are blocked - Public IPs (8.8.8.8, 1.1.1.1) are allowed - IPv4-mapped IPv6 unwrapping logic works correctly - Metadata endpoints and .internal/.local hostnames are blocked - Normal hostnames (github.com, clawhub.dev) are allowed Also documents a known gap: url::Url::host_str() returns bracketed IPv6 addresses that std::net::IpAddr cannot parse, so IPv4-mapped IPv6 URLs currently bypass IP-based checks in validate_fetch_url. Co-Authored-By: Claude Opus 4.6 * refactor(testing): extract TestGatewayBuilder to eliminate gateway test duplication Both ws_gateway_integration.rs and openai_compat_integration.rs manually constructed GatewayState with 19+ fields. Extracted to a shared builder in src/channels/web/test_helpers.rs that provides sensible defaults and lets tests override only what they need. Co-Authored-By: Claude Opus 4.6 * docs: add implementation plans for testing batches 1 and 2 Co-Authored-By: Claude Opus 4.6 * fix(security): close IPv6 SSRF bypass in validate_fetch_url validate_fetch_url used host_str() which returns bracketed IPv6 (e.g. "[::ffff:7f00:1]") that IpAddr::parse() cannot handle, silently skipping IP-based SSRF checks for all IPv6 URLs. Switch to url::Host enum matching to extract proper IpAddr values without string parsing. IPv4-mapped IPv6 addresses like ::ffff:127.0.0.1 are now correctly unwrapped and blocked. Co-Authored-By: Claude Opus 4.6 * test(skills): add activation criteria limits enforcement tests Adds test_activation_criteria_enforce_limits to verify that enforce_limits() correctly trims excess patterns (>5), keywords (>20), and tags (>10), and filters out short keywords/tags (<3 chars). Co-Authored-By: Claude Opus 4.6 * test(wasm): add security regression tests for WASM tool loader Add 6 tests covering: tool name path separator rejection, empty name rejection, nonexistent file handling, invalid WASM bytes rejection, dotfile discovery behavior, and subdirectory non-recursion. Co-Authored-By: Claude Opus 4.6 * refactor: address PR review feedback - Remove plan files from repo (ilblackdragon review) - Replace CLAUDE.md test tier rules with pointer to check-boundaries.sh - Add Check 4 to check-boundaries.sh: enforces integration tests are gated behind the 'integration' feature flag Co-Authored-By: Claude Opus 4.6 * ci: add try_connect silent-skip pattern check to check-boundaries.sh Check 5 catches try_connect() and similar silent-skip patterns in integration tests. Tests should use feature gates to fail loudly when prerequisites are missing, not silently return. Co-Authored-By: Claude Opus 4.6 * fix(security): harden skill fetch SSRF checks * fix(scripts): use bash arrays in check-boundaries.sh tier violation check Refactor Check 4 in check-boundaries.sh to use bash arrays and printf instead of string concatenation with echo -e. This is more robust with special characters in filenames and avoids portability concerns with echo -e. [skip-regression-check] Co-Authored-By: Claude Opus 4.6 --------- Co-authored-by: Claude Opus 4.6 --- .github/workflows/test.yml | 2 +- CLAUDE.md | 10 + scripts/check-boundaries.sh | 223 +++++++++++++++ src/channels/manager.rs | 103 +++++++ src/channels/web/mod.rs | 7 + src/channels/web/test_helpers.rs | 104 +++++++ src/skills/mod.rs | 68 +++++ src/testing.rs | 215 +++++++++++++- src/tools/builtin/skill_tools.rs | 433 +++++++++++++++++++++++++++-- src/tools/wasm/loader.rs | 158 +++++++++++ src/workspace/search.rs | 169 +++++++++++ tests/heartbeat_integration.rs | 2 +- tests/openai_compat_integration.rs | 70 +---- tests/workspace_integration.rs | 44 +-- tests/ws_gateway_integration.rs | 37 +-- 15 files changed, 1479 insertions(+), 166 deletions(-) create mode 100755 scripts/check-boundaries.sh create mode 100644 src/channels/web/test_helpers.rs diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 18c4269c..7c2f564c 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -14,7 +14,7 @@ jobs: matrix: include: - name: all-features - flags: "--all-features" + flags: "--features postgres,libsql,html-to-markdown" - name: default flags: "" - name: libsql-only diff --git a/CLAUDE.md b/CLAUDE.md index 4b8b89b4..c06c8537 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -45,6 +45,16 @@ cargo test test_name RUST_LOG=ironclaw=debug cargo run ``` +### Test Tiers + +| Tier | Command | What runs | External deps | +|------|---------|-----------|---------------| +| Unit | `cargo test` | All `mod tests` + self-contained integration tests | None | +| Integration | `cargo test --features integration` | + PostgreSQL-dependent tests | Running PostgreSQL | +| Live | `cargo test --features integration -- --ignored` | + LLM-dependent tests | PostgreSQL + LLM API keys | + +Run `bash scripts/check-boundaries.sh` to verify test tier gating and other architecture rules. + ## Project Structure ``` diff --git a/scripts/check-boundaries.sh b/scripts/check-boundaries.sh new file mode 100755 index 00000000..1fc072f6 --- /dev/null +++ b/scripts/check-boundaries.sh @@ -0,0 +1,223 @@ +#!/usr/bin/env bash +# Architecture boundary checks for IronClaw. +# Run as: bash scripts/check-boundaries.sh +# Returns non-zero if hard violations are found. + +set -euo pipefail + +REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +cd "$REPO_ROOT" + +violations=0 + +echo "=== Architecture Boundary Checks ===" +echo + +# -------------------------------------------------------------------------- +# Check 1: Direct database driver usage outside the db layer +# -------------------------------------------------------------------------- +# tokio_postgres:: and libsql:: types should only appear in: +# - src/db/ (the database abstraction layer) +# - src/workspace/repository.rs (workspace's own DB layer) +# - src/error.rs (needs From impls for driver error types) +# - src/app.rs (bootstraps/initialises the database) +# - src/testing.rs (test infrastructure) +# - src/cli/ (CLI commands that bootstrap DB connections) +# - src/setup/ (onboarding wizard bootstraps DB) +# - src/main.rs (entry point) +# +# Everything else is a boundary violation -- those modules should go through +# the Database trait, not touch driver types directly. +# -------------------------------------------------------------------------- + +echo "--- Check 1: Direct database driver usage outside db layer ---" + +results=$(grep -rn 'tokio_postgres::\|libsql::' src/ \ + --include='*.rs' \ + | grep -v 'src/db/' \ + | grep -v 'src/workspace/repository.rs' \ + | grep -v 'src/error.rs' \ + | grep -v 'src/app.rs' \ + | grep -v 'src/testing.rs' \ + | grep -v 'src/cli/' \ + | grep -v 'src/setup/' \ + | grep -v 'src/main.rs' \ + | grep -v '^\s*//' \ + | grep -v '//.*tokio_postgres\|//.*libsql' \ + || true) + +if [ -n "$results" ]; then + echo "VIOLATION: Direct database driver usage found outside db layer:" + echo "$results" + echo + count=$(echo "$results" | wc -l | tr -d ' ') + echo "($count occurrence(s) -- these modules should use the Database trait)" + violations=$((violations + 1)) +else + echo "OK" +fi +echo + +# -------------------------------------------------------------------------- +# Check 2: .unwrap() / .expect() in production code (heuristic) +# -------------------------------------------------------------------------- +# We cannot perfectly distinguish test vs production code with grep alone +# (test modules span many lines). Instead we: +# 1. Exclude files that are entirely test infrastructure +# 2. Exclude lines that are clearly in test code (assert, #[test], etc.) +# 3. Report a per-file summary so reviewers can focus on the worst files +# +# This is a WARNING, not a hard violation. +# -------------------------------------------------------------------------- + +echo "--- Check 2: .unwrap() / .expect() in production code ---" + +# Collect raw matches excluding obvious test-only files and lines +raw_results=$(grep -rn '\.unwrap()\|\.expect(' src/ \ + --include='*.rs' \ + | grep -v 'src/main.rs' \ + | grep -v 'src/testing.rs' \ + | grep -v 'src/setup/' \ + || true) + +if [ -n "$raw_results" ]; then + total=$(echo "$raw_results" | wc -l | tr -d ' ') + echo "WARNING: ~$total .unwrap()/.expect() calls found in src/ (excluding main/testing/setup)." + echo "Many are in test modules; a per-file breakdown helps triage:" + echo + # Show per-file counts, sorted by count descending, top 15 + file_counts=$(echo "$raw_results" | cut -d: -f1 | sort | uniq -c | sort -rn) + echo "$file_counts" | head -15 + fc_total=$(echo "$file_counts" | wc -l | tr -d ' ') + if [ "$fc_total" -gt 15 ]; then + echo " ... and $((fc_total - 15)) more files" + fi + echo + echo "(This is a warning for gradual cleanup, not a blocking violation.)" + echo "(Many of these are inside #[cfg(test)] modules which is acceptable.)" +else + echo "OK" +fi +echo + +# -------------------------------------------------------------------------- +# Check 3: std::env::var reads outside config/bootstrap layers +# -------------------------------------------------------------------------- +# Sensitive values should come through Config or the secrets module. +# Direct std::env::var / env::var() reads are allowed in: +# - src/config/ (the config layer itself) +# - src/main.rs (entry point) +# - src/setup/ (onboarding wizard) +# - src/testing.rs (test infrastructure) +# - src/cli/ (CLI commands that read env for bootstrap) +# - src/bootstrap.rs (bootstrap logic) +# -------------------------------------------------------------------------- + +echo "--- Check 3: Direct env var reads outside config layer ---" + +results=$(grep -rn 'std::env::var\|env::var(' src/ \ + --include='*.rs' \ + | grep -v 'src/config/' \ + | grep -v 'src/main.rs' \ + | grep -v 'src/setup/' \ + | grep -v 'src/testing.rs' \ + | grep -v 'src/cli/' \ + | grep -v 'src/bootstrap.rs' \ + | grep -v '#\[cfg(test)\]' \ + | grep -v '#\[test\]' \ + | grep -v 'mod tests' \ + | grep -v 'fn test_' \ + | grep -v '//.*env::var' \ + || true) + +if [ -n "$results" ]; then + count=$(echo "$results" | wc -l | tr -d ' ') + echo "WARNING: Direct env var reads found outside config layer ($count occurrences):" + echo "$results" + echo + echo "(Review these -- secrets/config should come through Config or the secrets module)" +else + echo "OK" +fi +echo + +# -------------------------------------------------------------------------- +# Check 4: Test tier gating — integration tests must use feature flags +# -------------------------------------------------------------------------- +# Files in tests/ that connect to PostgreSQL or use DATABASE_URL must be +# gated behind #![cfg(all(feature = "postgres", feature = "integration"))]. +# This ensures `cargo test` (no flags) never requires external services. +# +# Heuristic: any test file referencing DATABASE_URL, connect(), PgPool, +# or tokio_postgres should have the cfg gate on the first few lines. +# -------------------------------------------------------------------------- + +echo "--- Check 4: Test tier gating for integration tests ---" + +tier_violations=() +for test_file in tests/*.rs; do + [ -f "$test_file" ] || continue + + # Check if the file actually connects to a database (imports DB types + # or calls pool/connect). Mere string references like "DATABASE_URL" + # in config tests don't count. + needs_gate=false + if grep -q 'PgPool\|tokio_postgres::\|create_pool\|\.connect(' "$test_file" 2>/dev/null; then + needs_gate=true + fi + + if [ "$needs_gate" = true ]; then + # Check first 5 lines for the cfg gate + if ! head -5 "$test_file" | grep -q 'cfg.*feature.*integration' 2>/dev/null; then + tier_violations+=(" $test_file: needs '#![cfg(all(feature = \"postgres\", feature = \"integration\"))]'") + fi + fi +done + +if [ ${#tier_violations[@]} -gt 0 ]; then + echo "VIOLATION: Integration tests missing feature gate:" + printf '%s\n' "${tier_violations[@]}" + echo + echo "(Tests requiring external services must be gated behind the 'integration' feature)" + violations=$((violations + 1)) +else + echo "OK" +fi +echo + +# -------------------------------------------------------------------------- +# Check 5: No silent test-skip patterns (try_connect, is_available, etc.) +# -------------------------------------------------------------------------- +# Tests must fail loudly when prerequisites are missing, not silently skip. +# The correct approach is feature-flag gating (#![cfg(feature = "integration")]). +# Patterns like try_connect().is_none() { return; } hide broken tests. +# -------------------------------------------------------------------------- + +echo "--- Check 5: No silent test-skip patterns ---" + +skip_results=$(grep -rn 'try_connect\|is_available.*return\|is_none.*return\|is_err.*return.*//.*skip' tests/ \ + --include='*.rs' \ + || true) + +if [ -n "$skip_results" ]; then + echo "VIOLATION: Silent test-skip patterns found (use feature gates instead):" + echo "$skip_results" + echo + violations=$((violations + 1)) +else + echo "OK" +fi +echo + +# -------------------------------------------------------------------------- +# Summary +# -------------------------------------------------------------------------- + +echo "=== Summary ===" +if [ "$violations" -gt 0 ]; then + echo "FAILED: $violations hard violation(s) found" + exit 1 +else + echo "PASSED: No hard violations found (review warnings above)" + exit 0 +fi diff --git a/src/channels/manager.rs b/src/channels/manager.rs index 710c09c4..50d72e69 100644 --- a/src/channels/manager.rs +++ b/src/channels/manager.rs @@ -235,3 +235,106 @@ impl Default for ChannelManager { Self::new() } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::channels::IncomingMessage; + use crate::testing::StubChannel; + use futures::StreamExt; + + #[tokio::test] + async fn test_add_and_start_all() { + let manager = ChannelManager::new(); + let (stub, sender) = StubChannel::new("test"); + + manager.add(Box::new(stub)).await; + + let mut stream = manager.start_all().await.expect("start_all failed"); + + // Inject a message through the stub + sender + .send(IncomingMessage::new("test", "user1", "hello")) + .await + .expect("send failed"); + + // Should appear in the merged stream + let msg = stream.next().await.expect("stream ended"); + assert_eq!(msg.content, "hello"); + assert_eq!(msg.channel, "test"); + } + + #[tokio::test] + async fn test_respond_routes_to_correct_channel() { + let manager = ChannelManager::new(); + let (stub, _sender) = StubChannel::new("alpha"); + + // Keep a reference for response inspection + let responses = stub.captured_responses_handle(); + manager.add(Box::new(stub)).await; + + let msg = IncomingMessage::new("alpha", "user1", "request"); + manager + .respond(&msg, OutgoingResponse::text("reply")) + .await + .expect("respond failed"); + + // Verify the stub captured the response + let captured = responses.lock().expect("poisoned"); + assert_eq!(captured.len(), 1); + assert_eq!(captured[0].1.content, "reply"); + } + + #[tokio::test] + async fn test_respond_unknown_channel_errors() { + let manager = ChannelManager::new(); + let msg = IncomingMessage::new("nonexistent", "user1", "test"); + let result = manager.respond(&msg, OutgoingResponse::text("hi")).await; + assert!(result.is_err()); + } + + #[tokio::test] + async fn test_health_check_all() { + let manager = ChannelManager::new(); + let (stub1, _) = StubChannel::new("healthy"); + let (stub2, _) = StubChannel::new("sick"); + stub2.set_healthy(false); + + manager.add(Box::new(stub1)).await; + manager.add(Box::new(stub2)).await; + + let results = manager.health_check_all().await; + assert!(results["healthy"].is_ok()); + assert!(results["sick"].is_err()); + } + + #[tokio::test] + async fn test_start_all_no_channels_errors() { + let manager = ChannelManager::new(); + let result = manager.start_all().await; + assert!(result.is_err()); + } + + #[tokio::test] + async fn test_injection_channel_merges() { + let manager = ChannelManager::new(); + let (stub, _sender) = StubChannel::new("real"); + manager.add(Box::new(stub)).await; + + let mut stream = manager.start_all().await.expect("start_all failed"); + + // Use the injection channel (simulating background task) + let inject_tx = manager.inject_sender(); + inject_tx + .send(IncomingMessage::new( + "injected", + "system", + "background alert", + )) + .await + .expect("inject failed"); + + let msg = stream.next().await.expect("stream ended"); + assert_eq!(msg.content, "background alert"); + } +} diff --git a/src/channels/web/mod.rs b/src/channels/web/mod.rs index 5152e551..57597af0 100644 --- a/src/channels/web/mod.rs +++ b/src/channels/web/mod.rs @@ -24,6 +24,13 @@ pub mod types; pub(crate) mod util; pub mod ws; +/// Test helpers for gateway integration tests. +/// +/// Always compiled (not behind `#[cfg(test)]`) so that integration tests in +/// `tests/` -- which import this crate as a regular dependency -- can use +/// [`TestGatewayBuilder`](test_helpers::TestGatewayBuilder). +pub mod test_helpers; + use std::net::SocketAddr; use std::sync::Arc; diff --git a/src/channels/web/test_helpers.rs b/src/channels/web/test_helpers.rs new file mode 100644 index 00000000..9248f8f4 --- /dev/null +++ b/src/channels/web/test_helpers.rs @@ -0,0 +1,104 @@ +//! Shared test utilities for gateway integration tests. +//! +//! This module is always compiled (not `#[cfg(test)]`) because integration tests +//! in `tests/` import the crate as a regular dependency and `cfg(test)` is only +//! set when compiling *this* crate's unit tests. + +use std::net::SocketAddr; +use std::sync::Arc; + +use tokio::sync::mpsc; + +use crate::channels::IncomingMessage; +use crate::channels::web::server::{GatewayState, RateLimiter, start_server}; +use crate::channels::web::sse::SseManager; +use crate::channels::web::ws::WsConnectionTracker; + +/// Builder for constructing a [`GatewayState`] with sensible test defaults. +/// +/// Every optional field defaults to `None` and can be overridden via builder +/// methods. Call [`build`](Self::build) to get the `Arc`, or +/// [`start`](Self::start) to also bind an Axum server on a random port. +pub struct TestGatewayBuilder { + msg_tx: Option>, + llm_provider: Option>, + user_id: String, +} + +impl Default for TestGatewayBuilder { + fn default() -> Self { + Self { + msg_tx: None, + llm_provider: None, + user_id: "test-user".to_string(), + } + } +} + +impl TestGatewayBuilder { + /// Create a new builder with all defaults. + pub fn new() -> Self { + Self::default() + } + + /// Set the agent message sender (the channel the gateway forwards + /// incoming chat messages to). + pub fn msg_tx(mut self, tx: mpsc::Sender) -> Self { + self.msg_tx = Some(tx); + self + } + + /// Set the LLM provider (needed for OpenAI-compatible API tests). + pub fn llm_provider(mut self, provider: Arc) -> Self { + self.llm_provider = Some(provider); + self + } + + /// Override the user ID (default: `"test-user"`). + pub fn user_id(mut self, id: impl Into) -> Self { + self.user_id = id.into(); + self + } + + /// Build the `Arc` without starting a server. + pub fn build(self) -> Arc { + Arc::new(GatewayState { + msg_tx: tokio::sync::RwLock::new(self.msg_tx), + sse: SseManager::new(), + workspace: None, + session_manager: None, + log_broadcaster: None, + log_level_handle: None, + extension_manager: None, + tool_registry: None, + store: None, + job_manager: None, + prompt_queue: None, + user_id: self.user_id, + shutdown_tx: tokio::sync::RwLock::new(None), + ws_tracker: Some(Arc::new(WsConnectionTracker::new())), + llm_provider: self.llm_provider, + skill_registry: None, + skill_catalog: None, + scheduler: None, + chat_rate_limiter: RateLimiter::new(30, 60), + registry_entries: Vec::new(), + cost_guard: None, + startup_time: std::time::Instant::now(), + }) + } + + /// Build the state and start a gateway server on `127.0.0.1:0` (random + /// port). Returns the bound address and the shared state. + pub async fn start( + self, + auth_token: &str, + ) -> Result<(SocketAddr, Arc), crate::error::ChannelError> { + let state = self.build(); + let addr: SocketAddr = "127.0.0.1:0" + .parse() + .expect("hard-coded address must parse"); + let bound = start_server(addr, state.clone(), auth_token.to_string()).await?; + Ok((bound, state)) + } +} diff --git a/src/skills/mod.rs b/src/skills/mod.rs index 78407812..87e449c2 100644 --- a/src/skills/mod.rs +++ b/src/skills/mod.rs @@ -383,6 +383,74 @@ mod tests { assert_eq!(criteria.tags, vec!["foo", "bar"]); } + #[test] + fn test_activation_criteria_enforce_limits() { + // Build criteria that exceed all limits: + // - 25 keywords (5 over the 20 cap), including some short ones + // - 8 patterns (3 over the 5 cap) + // - 15 tags (5 over the 10 cap), including some short ones + let mut keywords: Vec = vec!["a".into(), "bb".into()]; // short, should be filtered + keywords.extend((0..25).map(|i| format!("keyword{}", i))); + + let patterns: Vec = (0..8).map(|i| format!("pattern{}", i)).collect(); + + let mut tags: Vec = vec!["x".into(), "ab".into()]; // short, should be filtered + tags.extend((0..15).map(|i| format!("tag{}", i))); + + let mut criteria = ActivationCriteria { + keywords, + patterns, + tags, + ..Default::default() + }; + + criteria.enforce_limits(); + + // Short keywords (<3 chars) filtered, then truncated to 20 + assert!( + !criteria + .keywords + .iter() + .any(|k| k.len() < MIN_KEYWORD_TAG_LENGTH), + "keywords shorter than {} chars should be filtered out", + MIN_KEYWORD_TAG_LENGTH + ); + assert_eq!( + criteria.keywords.len(), + MAX_KEYWORDS_PER_SKILL, + "keywords should be capped at {}", + MAX_KEYWORDS_PER_SKILL + ); + + // Patterns truncated to 5 (no length filter on patterns) + assert_eq!( + criteria.patterns.len(), + MAX_PATTERNS_PER_SKILL, + "patterns should be capped at {}", + MAX_PATTERNS_PER_SKILL + ); + // Verify the retained patterns are the first 5 + for i in 0..MAX_PATTERNS_PER_SKILL { + assert_eq!(criteria.patterns[i], format!("pattern{}", i)); + } + + // Short tags (<3 chars) filtered, then truncated to 10 + assert!( + !criteria + .tags + .iter() + .any(|t| t.len() < MIN_KEYWORD_TAG_LENGTH), + "tags shorter than {} chars should be filtered out", + MIN_KEYWORD_TAG_LENGTH + ); + assert_eq!( + criteria.tags.len(), + MAX_TAGS_PER_SKILL, + "tags should be capped at {}", + MAX_TAGS_PER_SKILL + ); + } + #[test] fn test_compile_patterns() { let patterns = vec![ diff --git a/src/testing.rs b/src/testing.rs index 7c36dc98..c62c2dcf 100644 --- a/src/testing.rs +++ b/src/testing.rs @@ -2,6 +2,7 @@ //! //! Provides: //! - [`StubLlm`]: A configurable LLM provider that returns a fixed response +//! - [`StubChannel`]: A configurable channel stub with message injection and response capture //! - [`TestHarnessBuilder`]: Builder for wiring `AgentDeps` with defaults //! - [`TestHarness`]: The assembled components ready for use in tests //! @@ -18,14 +19,19 @@ //! ``` use std::sync::Arc; +use std::sync::Mutex; use std::sync::atomic::{AtomicBool, AtomicU32, Ordering}; use async_trait::async_trait; use rust_decimal::Decimal; +use tokio::sync::mpsc; use crate::agent::AgentDeps; +use crate::channels::{ + Channel, ChannelManager, IncomingMessage, MessageStream, OutgoingResponse, StatusUpdate, +}; use crate::db::Database; -use crate::error::LlmError; +use crate::error::{ChannelError, LlmError}; use crate::llm::{ CompletionRequest, CompletionResponse, FinishReason, LlmProvider, ToolCompletionRequest, ToolCompletionResponse, @@ -189,12 +195,138 @@ impl LlmProvider for StubLlm { } } +/// A configurable channel stub for tests. +/// +/// Supports: +/// - Message injection via the returned `mpsc::Sender` +/// - Response capture for assertion +/// - Status update capture +/// - Configurable health check failure +/// +/// # Usage +/// +/// ```rust,no_run +/// let (channel, sender) = StubChannel::new("test"); +/// sender.send(IncomingMessage::new("test", "user1", "hello")).await.unwrap(); +/// // ... run agent logic that calls channel.respond() ... +/// let responses = channel.captured_responses(); +/// ``` +pub struct StubChannel { + name: String, + rx: tokio::sync::Mutex>>, + responses: Arc>>, + statuses: Arc>>, + healthy: AtomicBool, +} + +impl StubChannel { + /// Create a new stub channel and its message sender. + /// + /// The sender is used by tests to inject messages into the channel's stream. + /// The channel captures all responses and status updates for later assertion. + pub fn new(name: impl Into) -> (Self, mpsc::Sender) { + let (tx, rx) = mpsc::channel(64); + let channel = Self { + name: name.into(), + rx: tokio::sync::Mutex::new(Some(rx)), + responses: Arc::new(Mutex::new(Vec::new())), + statuses: Arc::new(Mutex::new(Vec::new())), + healthy: AtomicBool::new(true), + }; + (channel, tx) + } + + /// Get all captured (message, response) pairs. + pub fn captured_responses(&self) -> Vec<(IncomingMessage, OutgoingResponse)> { + self.responses.lock().expect("poisoned").clone() + } + + /// Get a shared handle to the response capture list. + /// + /// Call this *before* moving the channel into a `ChannelManager`, + /// since `add()` takes ownership. + pub fn captured_responses_handle( + &self, + ) -> Arc>> { + Arc::clone(&self.responses) + } + + /// Get all captured status updates. + pub fn captured_statuses(&self) -> Vec { + self.statuses.lock().expect("poisoned").clone() + } + + /// Get a shared handle to the status capture list. + pub fn captured_statuses_handle(&self) -> Arc>> { + Arc::clone(&self.statuses) + } + + /// Set whether `health_check()` succeeds or fails. + pub fn set_healthy(&self, healthy: bool) { + self.healthy.store(healthy, Ordering::Relaxed); + } +} + +#[async_trait] +impl Channel for StubChannel { + fn name(&self) -> &str { + &self.name + } + + async fn start(&self) -> Result { + let rx = self + .rx + .lock() + .await + .take() + .ok_or_else(|| ChannelError::StartupFailed { + name: self.name.clone(), + reason: "start() already called".to_string(), + })?; + let stream = tokio_stream::wrappers::ReceiverStream::new(rx); + Ok(Box::pin(stream)) + } + + async fn respond( + &self, + msg: &IncomingMessage, + response: OutgoingResponse, + ) -> Result<(), ChannelError> { + self.responses + .lock() + .expect("poisoned") + .push((msg.clone(), response)); + Ok(()) + } + + async fn send_status( + &self, + status: StatusUpdate, + _metadata: &serde_json::Value, + ) -> Result<(), ChannelError> { + self.statuses.lock().expect("poisoned").push(status); + Ok(()) + } + + async fn health_check(&self) -> Result<(), ChannelError> { + if self.healthy.load(Ordering::Relaxed) { + Ok(()) + } else { + Err(ChannelError::HealthCheckFailed { + name: self.name.clone(), + }) + } + } +} + /// Assembled test components. pub struct TestHarness { /// The agent dependencies, ready for use. pub deps: AgentDeps, /// Direct reference to the database (as `Arc`). pub db: Arc, + /// Stub channel sender + manager, present if `with_stub_channel()` was called. + pub channel: Option<(mpsc::Sender, ChannelManager)>, /// Temp directory guard — keeps the test database alive. Dropped /// automatically when the harness goes out of scope. #[cfg(feature = "libsql")] @@ -214,6 +346,7 @@ pub struct TestHarnessBuilder { db: Option>, llm: Option>, tools: Option>, + stub_channel: bool, } impl TestHarnessBuilder { @@ -223,6 +356,7 @@ impl TestHarnessBuilder { db: None, llm: None, tools: None, + stub_channel: false, } } @@ -244,6 +378,15 @@ impl TestHarnessBuilder { self } + /// Include a `StubChannel` wired into a `ChannelManager`. + /// + /// The harness will expose the sender (for injecting messages) and + /// the manager (for routing responses) via [`TestHarness::channel`]. + pub fn with_stub_channel(mut self) -> Self { + self.stub_channel = true; + self + } + /// Build the harness with defaults applied. #[cfg(feature = "libsql")] pub async fn build(self) -> TestHarness { @@ -280,6 +423,15 @@ impl TestHarnessBuilder { max_actions_per_hour: None, })); + let channel = if self.stub_channel { + let (stub, sender) = StubChannel::new("stub"); + let manager = ChannelManager::new(); + manager.add(Box::new(stub)).await; + Some((sender, manager)) + } else { + None + }; + let deps = AgentDeps { store: Some(Arc::clone(&db)), llm, @@ -300,6 +452,7 @@ impl TestHarnessBuilder { TestHarness { deps, db, + channel, _temp_dir: temp_dir, } } @@ -653,6 +806,48 @@ mod tests { assert_eq!(response.finish_reason, FinishReason::Stop); } + #[tokio::test] + async fn test_stub_channel_inject_and_capture() { + use futures::StreamExt; + + let (channel, sender) = StubChannel::new("test-channel"); + + // Start the channel to get the message stream + let mut stream = channel.start().await.expect("start failed"); + + // Inject a message + sender + .send(IncomingMessage::new("test-channel", "user1", "hello")) + .await + .expect("send failed"); + + // Read it from the stream + let msg = stream.next().await.expect("stream ended"); + assert_eq!(msg.content, "hello"); + assert_eq!(msg.user_id, "user1"); + assert_eq!(msg.channel, "test-channel"); + + // Send a response and verify it was captured + let response = OutgoingResponse::text("world"); + channel + .respond(&msg, response) + .await + .expect("respond failed"); + + let captured = channel.captured_responses(); + assert_eq!(captured.len(), 1); + assert_eq!(captured[0].1.content, "world"); + } + + #[tokio::test] + async fn test_stub_channel_health_check() { + let (channel, _sender) = StubChannel::new("healthy"); + channel.health_check().await.expect("health check failed"); + + channel.set_healthy(false); + assert!(channel.health_check().await.is_err()); + } + // === Database CRUD coverage for untested trait methods === #[cfg(feature = "libsql")] @@ -705,6 +900,24 @@ mod tests { assert!(!deleted); } + #[tokio::test] + async fn test_harness_with_channel() { + let harness = TestHarnessBuilder::new().with_stub_channel().build().await; + + let (sender, channel_manager) = + harness.channel.as_ref().expect("channel should be present"); + + // Inject a message via sender + sender + .send(IncomingMessage::new("stub", "user1", "test message")) + .await + .expect("send failed"); + + // Verify channel is registered in the manager + let names = channel_manager.channel_names().await; + assert!(names.contains(&"stub".to_string())); + } + #[cfg(feature = "libsql")] #[tokio::test] async fn test_settings_bulk_operations() { diff --git a/src/tools/builtin/skill_tools.rs b/src/tools/builtin/skill_tools.rs index 65948d27..84c889ae 100644 --- a/src/tools/builtin/skill_tools.rs +++ b/src/tools/builtin/skill_tools.rs @@ -380,8 +380,8 @@ impl Tool for SkillInstallTool { /// - Non-HTTPS URLs (except in tests) /// - URLs pointing to private, loopback, or link-local IP addresses /// - URLs without a host -pub fn validate_fetch_url(url_str: &str) -> Result<(), ToolError> { - let parsed = url::Url::parse(url_str) +pub fn validate_fetch_url(url_str: &str) -> Result { + let parsed = reqwest::Url::parse(url_str) .map_err(|e| ToolError::ExecutionFailed(format!("Invalid URL '{}': {}", url_str, e)))?; // Require HTTPS @@ -393,30 +393,20 @@ pub fn validate_fetch_url(url_str: &str) -> Result<(), ToolError> { } let host = parsed - .host_str() + .host() .ok_or_else(|| ToolError::ExecutionFailed("URL has no host".to_string()))?; // Check if host is an IP address and reject private ranges. + // Use reqwest::Url host variants to get proper IpAddr values -- host_str() + // returns bracketed IPv6 (e.g. "[::1]") which IpAddr cannot parse. // Unwrap IPv4-mapped IPv6 addresses (e.g. ::ffff:192.168.1.1) to catch // SSRF bypasses that encode private IPv4 addresses as IPv6. - if let Ok(raw_ip) = host.parse::() { - let ip = match raw_ip { - std::net::IpAddr::V6(v6) => v6 - .to_ipv4_mapped() - .map(std::net::IpAddr::V4) - .unwrap_or(std::net::IpAddr::V6(v6)), - other => other, - }; - if ip.is_loopback() || ip.is_unspecified() || is_private_ip(&ip) || is_link_local_ip(&ip) { - return Err(ToolError::ExecutionFailed(format!( - "URL points to a private/loopback/link-local address: {}", - host - ))); - } + if let Some(ip) = host_ip_addr(&host) { + validate_fetch_ip(&ip, &host.to_string())?; } - // Reject common internal hostnames - let host_lower = host.to_lowercase(); + // Reject common internal hostnames, including FQDN forms with a trailing dot. + let host_lower = normalize_domain(host.to_string().as_str()).to_lowercase(); if host_lower == "localhost" || host_lower == "metadata.google.internal" || host_lower.ends_with(".internal") @@ -428,9 +418,100 @@ pub fn validate_fetch_url(url_str: &str) -> Result<(), ToolError> { ))); } + Ok(parsed) +} + +fn host_ip_addr(host: &url::Host<&str>) -> Option { + match host { + url::Host::Ipv4(v4) => Some(std::net::IpAddr::V4(*v4)), + url::Host::Ipv6(v6) => Some(normalize_ip(std::net::IpAddr::V6(*v6))), + url::Host::Domain(_) => None, + } +} + +fn normalize_ip(ip: std::net::IpAddr) -> std::net::IpAddr { + match ip { + std::net::IpAddr::V6(v6) => v6 + .to_ipv4_mapped() + .map(std::net::IpAddr::V4) + .unwrap_or(std::net::IpAddr::V6(v6)), + other => other, + } +} + +fn validate_fetch_ip(ip: &std::net::IpAddr, display_host: &str) -> Result<(), ToolError> { + if ip.is_loopback() || ip.is_unspecified() || is_private_ip(ip) || is_link_local_ip(ip) { + return Err(ToolError::ExecutionFailed(format!( + "URL points to a private/loopback/link-local address: {}", + display_host + ))); + } + Ok(()) } +fn normalize_domain(host: &str) -> &str { + host.trim_end_matches('.') +} + +fn validate_resolved_addrs(host: &str, addrs: &[std::net::SocketAddr]) -> Result<(), ToolError> { + if addrs.is_empty() { + return Err(ToolError::ExecutionFailed(format!( + "DNS resolution returned no addresses for {}", + host + ))); + } + + for addr in addrs { + let ip = normalize_ip(addr.ip()); + validate_fetch_ip(&ip, host)?; + } + + Ok(()) +} + +fn build_fetch_client_builder() -> reqwest::ClientBuilder { + reqwest::Client::builder() + .timeout(std::time::Duration::from_secs(15)) + .user_agent("ironclaw/0.1") + .redirect(reqwest::redirect::Policy::none()) +} + +async fn build_safe_fetch_client(parsed: &reqwest::Url) -> Result { + let host = parsed + .host() + .ok_or_else(|| ToolError::ExecutionFailed("URL has no host".to_string()))?; + + match host { + url::Host::Ipv4(_) | url::Host::Ipv6(_) => build_fetch_client_builder() + .build() + .map_err(|e| ToolError::ExecutionFailed(format!("HTTP client error: {}", e))), + url::Host::Domain(domain) => { + let lookup_host = normalize_domain(domain); + let port = parsed + .port_or_known_default() + .ok_or_else(|| ToolError::ExecutionFailed("URL has no valid port".to_string()))?; + + let addrs: Vec = tokio::net::lookup_host((lookup_host, port)) + .await + .map_err(|e| { + ToolError::ExecutionFailed(format!( + "DNS resolution failed for {}: {}", + lookup_host, e + )) + })? + .collect(); + + validate_resolved_addrs(domain, &addrs)?; + + build_fetch_client_builder() + .resolve_to_addrs(domain, &addrs) + .build() + .map_err(|e| ToolError::ExecutionFailed(format!("HTTP client error: {}", e))) + } + } +} + fn is_private_ip(ip: &std::net::IpAddr) -> bool { match ip { std::net::IpAddr::V4(v4) => { @@ -463,16 +544,10 @@ fn is_link_local_ip(ip: &std::net::IpAddr) -> bool { /// `PK\x03\x04` magic bytes) and extracts `SKILL.md` automatically. Plain /// text responses are returned as-is. pub async fn fetch_skill_content(url: &str) -> Result { - validate_fetch_url(url)?; + let parsed = validate_fetch_url(url)?; + let client = build_safe_fetch_client(&parsed).await?; - let client = reqwest::Client::builder() - .timeout(std::time::Duration::from_secs(15)) - .user_agent("ironclaw/0.1") - .redirect(reqwest::redirect::Policy::none()) - .build() - .map_err(|e| ToolError::ExecutionFailed(format!("HTTP client error: {}", e)))?; - - let response = client.get(url).send().await.map_err(|e| { + let response = client.get(parsed.clone()).send().await.map_err(|e| { ToolError::ExecutionFailed(format!("Failed to fetch skill from {}: {}", url, e)) })?; @@ -797,6 +872,12 @@ mod tests { assert!(err.to_string().contains("internal hostname")); } + #[test] + fn test_validate_fetch_url_rejects_localhost_fqdn() { + let err = super::validate_fetch_url("https://localhost./skill.md").unwrap_err(); + assert!(err.to_string().contains("internal hostname")); + } + #[test] fn test_validate_fetch_url_rejects_metadata_endpoint() { let err = @@ -817,6 +898,41 @@ mod tests { assert!(err.to_string().contains("Only HTTPS")); } + #[test] + fn test_validate_fetch_url_rejects_ipv4_mapped_ipv6_loopback() { + let err = super::validate_fetch_url("https://[::ffff:127.0.0.1]/skill.md").unwrap_err(); + assert!(err.to_string().contains("private") || err.to_string().contains("loopback")); + } + + #[test] + fn test_validate_fetch_url_rejects_ipv6_loopback() { + let err = super::validate_fetch_url("https://[::1]/skill.md").unwrap_err(); + assert!(err.to_string().contains("private") || err.to_string().contains("loopback")); + } + + #[test] + fn test_validate_resolved_addrs_rejects_loopback_hostname() { + let addrs = vec![ + "127.0.0.1:443".parse::().unwrap(), + "[::1]:443".parse::().unwrap(), + ]; + + let err = super::validate_resolved_addrs("example.com", &addrs).unwrap_err(); + assert!(err.to_string().contains("private") || err.to_string().contains("loopback")); + } + + #[test] + fn test_validate_resolved_addrs_allows_public_hostname() { + let addrs = vec![ + "8.8.8.8:443".parse::().unwrap(), + "[2606:4700:4700::1111]:443" + .parse::() + .unwrap(), + ]; + + assert!(super::validate_resolved_addrs("example.com", &addrs).is_ok()); + } + #[test] fn test_extract_skill_from_zip_deflate() { // Build a real ZIP with flate2 + manual header construction. @@ -890,4 +1006,265 @@ mod tests { let err = super::extract_skill_from_zip(&zip).unwrap_err(); assert!(err.to_string().contains("does not contain SKILL.md")); } + + // ── ZIP extraction security regression tests ──────────────────────── + + /// Helper: build a minimal ZIP local file header with Store compression. + fn build_zip_entry_store(file_name: &str, content: &[u8]) -> Vec { + let mut zip = Vec::new(); + zip.extend_from_slice(&[0x50, 0x4B, 0x03, 0x04]); // signature + zip.extend_from_slice(&[0x0A, 0x00]); // version needed (1.0) + zip.extend_from_slice(&[0x00, 0x00]); // flags + zip.extend_from_slice(&[0x00, 0x00]); // compression: store (0) + zip.extend_from_slice(&[0x00, 0x00, 0x00, 0x00]); // mod time/date + zip.extend_from_slice(&[0x00, 0x00, 0x00, 0x00]); // crc32 + zip.extend_from_slice(&(content.len() as u32).to_le_bytes()); // compressed size + zip.extend_from_slice(&(content.len() as u32).to_le_bytes()); // uncompressed size + zip.extend_from_slice(&(file_name.len() as u16).to_le_bytes()); // filename length + zip.extend_from_slice(&0u16.to_le_bytes()); // extra field length + zip.extend_from_slice(file_name.as_bytes()); + zip.extend_from_slice(content); + zip + } + + #[test] + fn test_zip_extract_valid_skill() { + let content = b"---\nname: hello\n---\n# Hello Skill\nDoes things.\n"; + let zip = build_zip_entry_store("SKILL.md", content); + let result = super::extract_skill_from_zip(&zip).unwrap(); + assert_eq!(result, std::str::from_utf8(content).unwrap()); + } + + #[test] + fn test_zip_extract_ignores_non_skill_entries() { + // ZIP with README.md and src/main.rs but no SKILL.md -- should error. + let mut zip = Vec::new(); + zip.extend_from_slice(&build_zip_entry_store("README.md", b"# Readme")); + zip.extend_from_slice(&build_zip_entry_store("src/main.rs", b"fn main() {}")); + + let err = super::extract_skill_from_zip(&zip).unwrap_err(); + assert!( + err.to_string().contains("does not contain SKILL.md"), + "Expected 'does not contain SKILL.md' error, got: {}", + err + ); + } + + #[test] + fn test_zip_extract_path_traversal_rejected() { + // An entry named "../../SKILL.md" must NOT match the exact "SKILL.md" check. + let content = b"---\nname: evil\n---\n# Malicious path traversal\n"; + let zip = build_zip_entry_store("../../SKILL.md", content); + + let err = super::extract_skill_from_zip(&zip).unwrap_err(); + assert!( + err.to_string().contains("does not contain SKILL.md"), + "Path traversal entry should not match SKILL.md, got: {}", + err + ); + } + + #[test] + fn test_zip_extract_nested_path_not_matched() { + // An entry named "subdir/SKILL.md" must NOT match the exact "SKILL.md" check. + let content = b"---\nname: nested\n---\n# Nested\n"; + let zip = build_zip_entry_store("subdir/SKILL.md", content); + + let err = super::extract_skill_from_zip(&zip).unwrap_err(); + assert!( + err.to_string().contains("does not contain SKILL.md"), + "Nested path should not match SKILL.md, got: {}", + err + ); + } + + #[test] + fn test_zip_extract_oversized_rejected() { + // Create a ZIP entry whose declared uncompressed_size exceeds MAX_DECOMPRESSED (1 MB). + let oversized_claim: u32 = 2 * 1024 * 1024; // 2 MB + let small_body = b"tiny"; + + let mut zip = Vec::new(); + zip.extend_from_slice(&[0x50, 0x4B, 0x03, 0x04]); // signature + zip.extend_from_slice(&[0x0A, 0x00]); // version needed + zip.extend_from_slice(&[0x00, 0x00]); // flags + zip.extend_from_slice(&[0x00, 0x00]); // compression: store + zip.extend_from_slice(&[0x00, 0x00, 0x00, 0x00]); // mod time/date + zip.extend_from_slice(&[0x00, 0x00, 0x00, 0x00]); // crc32 + zip.extend_from_slice(&(small_body.len() as u32).to_le_bytes()); // compressed size (actual) + zip.extend_from_slice(&oversized_claim.to_le_bytes()); // uncompressed size (forged) + zip.extend_from_slice(&8u16.to_le_bytes()); // filename length + zip.extend_from_slice(&0u16.to_le_bytes()); // extra field length + zip.extend_from_slice(b"SKILL.md"); + zip.extend_from_slice(small_body); + + let err = super::extract_skill_from_zip(&zip).unwrap_err(); + assert!( + err.to_string().contains("too large"), + "Oversized entry should be rejected, got: {}", + err + ); + } + + // ── SSRF prevention regression tests ──────────────────────────────── + + #[test] + fn test_is_private_ip_blocks_loopback() { + let loopback: std::net::IpAddr = "127.0.0.1".parse().unwrap(); + // is_private_ip checks v4.is_private() which does NOT include loopback, + // but validate_fetch_url checks is_loopback() separately. Test the full flow. + assert!(loopback.is_loopback()); + // Also verify via validate_fetch_url + assert!(super::validate_fetch_url("https://127.0.0.1/skill.md").is_err()); + } + + #[test] + fn test_is_private_ip_blocks_private_ranges() { + let cases: Vec<(&str, bool)> = vec![ + ("10.0.0.1", true), + ("10.255.255.255", true), + ("172.16.0.1", true), + ("172.31.255.255", true), + ("192.168.1.1", true), + ("192.168.0.0", true), + ]; + for (ip_str, expect_private) in cases { + let ip: std::net::IpAddr = ip_str.parse().unwrap(); + assert_eq!( + super::is_private_ip(&ip), + expect_private, + "Expected is_private_ip({}) = {}", + ip_str, + expect_private + ); + } + } + + #[test] + fn test_is_private_ip_blocks_link_local() { + // 169.254.0.0/16 range (link-local) + let cases = vec!["169.254.1.1", "169.254.0.1", "169.254.255.255"]; + for ip_str in cases { + let ip: std::net::IpAddr = ip_str.parse().unwrap(); + // is_private_ip includes v4.is_link_local() + assert!( + super::is_private_ip(&ip), + "Expected is_private_ip({}) = true (link-local)", + ip_str + ); + } + } + + #[test] + fn test_is_private_ip_allows_public() { + let public_ips = vec!["8.8.8.8", "1.1.1.1", "93.184.216.34", "151.101.1.67"]; + for ip_str in public_ips { + let ip: std::net::IpAddr = ip_str.parse().unwrap(); + assert!( + !super::is_private_ip(&ip), + "Expected is_private_ip({}) = false (public IP)", + ip_str + ); + assert!(!ip.is_loopback(), "Expected {} is not loopback", ip_str); + } + } + + #[test] + fn test_is_private_ip_blocks_ipv4_mapped_ipv6() { + // Test the IPv4-mapped unwrapping logic end-to-end through + // validate_fetch_url. IPv6 URLs like https://[::ffff:127.0.0.1]/path + // must be correctly detected as private/loopback. + + // ::ffff:127.0.0.1 mapped -> 127.0.0.1 (loopback) -- must be blocked + let err = super::validate_fetch_url("https://[::ffff:127.0.0.1]/skill.md").unwrap_err(); + assert!( + err.to_string().contains("private") || err.to_string().contains("loopback"), + "IPv4-mapped loopback should be blocked, got: {}", + err + ); + + // ::ffff:192.168.1.1 mapped -> 192.168.1.1 (private) -- must be blocked + let err = super::validate_fetch_url("https://[::ffff:192.168.1.1]/skill.md").unwrap_err(); + assert!( + err.to_string().contains("private") || err.to_string().contains("loopback"), + "IPv4-mapped private should be blocked, got: {}", + err + ); + + // ::ffff:10.0.0.1 mapped -> 10.0.0.1 (private) -- must be blocked + let err = super::validate_fetch_url("https://[::ffff:10.0.0.1]/skill.md").unwrap_err(); + assert!( + err.to_string().contains("private") || err.to_string().contains("loopback"), + "IPv4-mapped 10.x should be blocked, got: {}", + err + ); + + // ::ffff:8.8.8.8 mapped -> 8.8.8.8 (public) -- must be allowed + assert!( + super::validate_fetch_url("https://[::ffff:8.8.8.8]/skill.md").is_ok(), + "IPv4-mapped public IP should be allowed" + ); + + // Pure IPv6 loopback ::1 -- must be blocked + let err = super::validate_fetch_url("https://[::1]/skill.md").unwrap_err(); + assert!( + err.to_string().contains("private") || err.to_string().contains("loopback"), + "IPv6 loopback should be blocked, got: {}", + err + ); + } + + #[test] + fn test_is_restricted_host_blocks_metadata() { + // Cloud metadata endpoint (AWS/GCP/Azure style) + let err = + super::validate_fetch_url("https://169.254.169.254/latest/meta-data/").unwrap_err(); + assert!( + err.to_string().contains("private") || err.to_string().contains("link-local"), + "Metadata IP should be blocked, got: {}", + err + ); + + // GCP metadata hostname + let err = + super::validate_fetch_url("https://metadata.google.internal/something").unwrap_err(); + assert!( + err.to_string().contains("internal hostname"), + "metadata.google.internal should be blocked, got: {}", + err + ); + + // Generic .internal domain + let err = super::validate_fetch_url("https://service.internal/api").unwrap_err(); + assert!( + err.to_string().contains("internal hostname"), + ".internal domains should be blocked, got: {}", + err + ); + + // .local domain + let err = super::validate_fetch_url("https://myhost.local/skill.md").unwrap_err(); + assert!( + err.to_string().contains("internal hostname"), + ".local domains should be blocked, got: {}", + err + ); + } + + #[test] + fn test_is_restricted_host_allows_normal() { + let allowed = vec![ + "https://github.com/repo/SKILL.md", + "https://clawhub.dev/api/v1/download?slug=foo", + "https://raw.githubusercontent.com/user/repo/main/SKILL.md", + "https://example.com/skills/deploy.md", + ]; + for url in allowed { + assert!( + super::validate_fetch_url(url).is_ok(), + "Expected validate_fetch_url({}) to succeed", + url + ); + } + } } diff --git a/src/tools/wasm/loader.rs b/src/tools/wasm/loader.rs index 7c87e568..ab94553e 100644 --- a/src/tools/wasm/loader.rs +++ b/src/tools/wasm/loader.rs @@ -919,4 +919,162 @@ mod tests { assert!(!config.client_id.is_empty()); assert!(config.client_secret.is_some()); } + + // --------------------------------------------------------------- + // Security regression tests + // --------------------------------------------------------------- + + use std::sync::Arc; + + use crate::tools::registry::ToolRegistry; + use crate::tools::wasm::{WasmRuntimeConfig, WasmToolRuntime}; + + /// Helper: create a WasmToolLoader backed by a real runtime + registry. + fn make_loader() -> super::WasmToolLoader { + let runtime = Arc::new( + WasmToolRuntime::new(WasmRuntimeConfig::for_testing()) + .expect("failed to create WASM runtime for test"), + ); + let registry = Arc::new(ToolRegistry::new()); + super::WasmToolLoader::new(runtime, registry) + } + + #[tokio::test] + async fn test_tool_name_rejects_path_separators() { + let dir = TempDir::new().unwrap(); + // Create a valid wasm file so the name check is the only failure path + let wasm_path = dir.path().join("dummy.wasm"); + std::fs::File::create(&wasm_path).unwrap(); + + let loader = make_loader(); + + for bad_name in &["../evil", "foo/bar", "foo\\bar"] { + let result = loader.load_from_files(bad_name, &wasm_path, None).await; + assert!( + result.is_err(), + "Expected error for name {:?}, got Ok", + bad_name + ); + let err = result.unwrap_err(); + assert!( + matches!(err, WasmLoadError::InvalidName(_)), + "Expected InvalidName for {:?}, got: {}", + bad_name, + err + ); + } + } + + #[tokio::test] + async fn test_tool_name_rejects_empty() { + let dir = TempDir::new().unwrap(); + let wasm_path = dir.path().join("dummy.wasm"); + std::fs::File::create(&wasm_path).unwrap(); + + let loader = make_loader(); + let result = loader.load_from_files("", &wasm_path, None).await; + + assert!(result.is_err(), "Expected error for empty name, got Ok"); + let err = result.unwrap_err(); + assert!( + matches!(err, WasmLoadError::InvalidName(_)), + "Expected InvalidName for empty string, got: {}", + err + ); + } + + #[tokio::test] + async fn test_load_nonexistent_wasm_file() { + let loader = make_loader(); + let bogus_path = std::path::PathBuf::from("/tmp/nonexistent_tool_12345.wasm"); + + let result = loader.load_from_files("bogus", &bogus_path, None).await; + assert!( + result.is_err(), + "Expected error for nonexistent file, got Ok" + ); + let err = result.unwrap_err(); + assert!( + matches!(err, WasmLoadError::WasmNotFound(_)), + "Expected WasmNotFound, got: {}", + err + ); + } + + #[tokio::test] + async fn test_load_invalid_wasm_bytes() { + let dir = TempDir::new().unwrap(); + let wasm_path = dir.path().join("invalid.wasm"); + + // Write random invalid bytes (not a valid WASM module) + let mut f = std::fs::File::create(&wasm_path).unwrap(); + f.write_all(b"this is not a valid wasm module at all") + .unwrap(); + + let loader = make_loader(); + let result = loader.load_from_files("invalid", &wasm_path, None).await; + + assert!( + result.is_err(), + "Expected error for invalid WASM bytes, got Ok" + ); + // The error should come from WASM compilation or registration, not name validation + let err = result.unwrap_err(); + assert!( + !matches!(err, WasmLoadError::InvalidName(_)), + "Got InvalidName instead of a compilation/registration error: {}", + err + ); + } + + #[tokio::test] + async fn test_discover_skips_dotfiles() { + let dir = TempDir::new().unwrap(); + + // Create a dotfile .wasm and a normal .wasm + std::fs::File::create(dir.path().join(".hidden.wasm")).unwrap(); + std::fs::File::create(dir.path().join("visible.wasm")).unwrap(); + + let tools = discover_tools(dir.path()).await.unwrap(); + + // The current implementation discovers ALL .wasm files including dotfiles. + // This test documents the current behavior: .hidden.wasm IS discovered + // with the stem ".hidden". A future hardening pass could add dotfile + // filtering, at which point this assertion should be updated. + assert!( + tools.contains_key("visible"), + "visible.wasm should be discovered" + ); + assert!( + tools.contains_key(".hidden"), + "dotfile .hidden.wasm is currently discovered (no dotfile filter yet)" + ); + assert_eq!(tools.len(), 2); + } + + #[tokio::test] + async fn test_discover_tools_ignores_subdirectories() { + let dir = TempDir::new().unwrap(); + + // Create a top-level wasm file + std::fs::File::create(dir.path().join("top_level.wasm")).unwrap(); + + // Create a subdirectory with a wasm file inside + let sub_dir = dir.path().join("subdir"); + std::fs::create_dir(&sub_dir).unwrap(); + std::fs::File::create(sub_dir.join("nested.wasm")).unwrap(); + + let tools = discover_tools(dir.path()).await.unwrap(); + + // Only top-level files should be discovered (read_dir is not recursive) + assert_eq!(tools.len(), 1, "Only top-level .wasm files should be found"); + assert!( + tools.contains_key("top_level"), + "top_level.wasm should be discovered" + ); + assert!( + !tools.contains_key("nested"), + "nested.wasm inside subdir should NOT be discovered" + ); + } } diff --git a/src/workspace/search.rs b/src/workspace/search.rs index 29e21c33..dff15298 100644 --- a/src/workspace/search.rs +++ b/src/workspace/search.rs @@ -458,4 +458,173 @@ mod tests { assert!(!vector_only.use_fts); assert!(vector_only.use_vector); } + + // --- Edge case tests --- + + #[test] + fn test_rrf_both_empty() { + let config = SearchConfig::default(); + let results = reciprocal_rank_fusion(Vec::new(), Vec::new(), &config); + assert!(results.is_empty()); + } + + #[test] + fn test_rrf_fts_only_no_vector() { + let config = SearchConfig::default().with_limit(10); + + let chunk1 = Uuid::new_v4(); + let chunk2 = Uuid::new_v4(); + let chunk3 = Uuid::new_v4(); + let doc = Uuid::new_v4(); + + let fts_results = vec![ + make_result(chunk1, doc, 1), + make_result(chunk2, doc, 2), + make_result(chunk3, doc, 3), + ]; + + let results = reciprocal_rank_fusion(fts_results, Vec::new(), &config); + + assert_eq!(results.len(), 3); + // All results should come from FTS only + assert!(results.iter().all(|r| r.from_fts())); + assert!(results.iter().all(|r| !r.from_vector())); + assert!(results.iter().all(|r| !r.is_hybrid())); + // Scores should be in descending order + for w in results.windows(2) { + assert!(w[0].score >= w[1].score); + } + } + + #[test] + fn test_rrf_vector_only_no_fts() { + let config = SearchConfig::default().with_limit(10); + + let chunk1 = Uuid::new_v4(); + let chunk2 = Uuid::new_v4(); + let chunk3 = Uuid::new_v4(); + let doc = Uuid::new_v4(); + + let vector_results = vec![ + make_result(chunk1, doc, 1), + make_result(chunk2, doc, 2), + make_result(chunk3, doc, 3), + ]; + + let results = reciprocal_rank_fusion(Vec::new(), vector_results, &config); + + assert_eq!(results.len(), 3); + // All results should come from vector only + assert!(results.iter().all(|r| r.from_vector())); + assert!(results.iter().all(|r| !r.from_fts())); + assert!(results.iter().all(|r| !r.is_hybrid())); + // Scores should be in descending order + for w in results.windows(2) { + assert!(w[0].score >= w[1].score); + } + } + + #[test] + fn test_rrf_duplicate_chunks_merged() { + let config = SearchConfig::default().with_limit(10); + + let shared_chunk = Uuid::new_v4(); + let fts_only_chunk = Uuid::new_v4(); + let vector_only_chunk = Uuid::new_v4(); + let doc = Uuid::new_v4(); + + // shared_chunk appears at rank 2 in FTS and rank 3 in vector + let fts_results = vec![ + make_result(fts_only_chunk, doc, 1), + make_result(shared_chunk, doc, 2), + ]; + let vector_results = vec![ + make_result(vector_only_chunk, doc, 1), + make_result(shared_chunk, doc, 3), + ]; + + let results = reciprocal_rank_fusion(fts_results, vector_results, &config); + + // Should have 3 unique chunks (not 4) + assert_eq!(results.len(), 3); + + // Find the shared chunk in results + let shared = results.iter().find(|r| r.chunk_id == shared_chunk).unwrap(); + assert!(shared.is_hybrid()); + assert_eq!(shared.fts_rank, Some(2)); + assert_eq!(shared.vector_rank, Some(3)); + + // The shared chunk's pre-normalization score is 1/(k+2) + 1/(k+3), + // which is higher than either single-method chunk at rank 1: 1/(k+1). + // After normalization the shared chunk should be the top result. + assert_eq!(results[0].chunk_id, shared_chunk); + } + + #[test] + fn test_rrf_limit_zero_returns_empty() { + let config = SearchConfig::default().with_limit(0); + + let doc = Uuid::new_v4(); + let fts_results = vec![ + make_result(Uuid::new_v4(), doc, 1), + make_result(Uuid::new_v4(), doc, 2), + ]; + + let results = reciprocal_rank_fusion(fts_results, Vec::new(), &config); + + assert!(results.is_empty()); + } + + #[test] + fn test_rrf_min_score_one_filters_all() { + // RRF scores are always < 1.0 before normalization (1/(k+rank) where k>=1, rank>=1). + // After normalization the top result gets score=1.0, so min_score=1.0 should + // keep only the single top result. To truly filter everything, we need + // min_score > 1.0 -- but with_min_score clamps to 1.0. + // With a single result: normalized score = 1.0, so it passes min_score=1.0. + // With multiple results: only the top (score=1.0) survives. + // To filter ALL results we need to ensure none reach 1.0 -- but normalization + // always makes the max = 1.0. So min_score=1.0 keeps exactly 1 result (the top). + // + // Verified: the retain check is `score >= min_score` and the top score + // is normalized to exactly 1.0, so one result survives. + let config = SearchConfig::default().with_limit(10).with_min_score(1.0); + + let doc = Uuid::new_v4(); + let fts_results = vec![ + make_result(Uuid::new_v4(), doc, 1), + make_result(Uuid::new_v4(), doc, 2), + make_result(Uuid::new_v4(), doc, 3), + ]; + + let results = reciprocal_rank_fusion(fts_results, Vec::new(), &config); + + // After normalization the top result has score 1.0, so exactly 1 survives + assert_eq!(results.len(), 1); + assert!((results[0].score - 1.0).abs() < 0.001); + } + + #[test] + fn test_search_config_fts_only() { + let config = SearchConfig::default().fts_only(); + + assert!(config.use_fts); + assert!(!config.use_vector); + // Other defaults should be preserved + assert_eq!(config.limit, 10); + assert_eq!(config.rrf_k, 60); + assert!((config.min_score - 0.0).abs() < f32::EPSILON); + } + + #[test] + fn test_search_config_vector_only() { + let config = SearchConfig::default().vector_only(); + + assert!(!config.use_fts); + assert!(config.use_vector); + // Other defaults should be preserved + assert_eq!(config.limit, 10); + assert_eq!(config.rrf_k, 60); + assert!((config.min_score - 0.0).abs() < f32::EPSILON); + } } diff --git a/tests/heartbeat_integration.rs b/tests/heartbeat_integration.rs index 227f59f9..f609b769 100644 --- a/tests/heartbeat_integration.rs +++ b/tests/heartbeat_integration.rs @@ -1,4 +1,4 @@ -#![cfg(feature = "postgres")] +#![cfg(all(feature = "postgres", feature = "integration"))] //! Heartbeat integration test. //! //! Exercises the heartbeat system in isolation: connects to the real diff --git a/tests/openai_compat_integration.rs b/tests/openai_compat_integration.rs index e70f895a..379978cb 100644 --- a/tests/openai_compat_integration.rs +++ b/tests/openai_compat_integration.rs @@ -9,9 +9,8 @@ use std::time::Duration; use async_trait::async_trait; use rust_decimal::Decimal; -use ironclaw::channels::web::server::{GatewayState, start_server}; -use ironclaw::channels::web::sse::SseManager; -use ironclaw::channels::web::ws::WsConnectionTracker; +use ironclaw::channels::web::server::GatewayState; +use ironclaw::channels::web::test_helpers::TestGatewayBuilder; use ironclaw::error::LlmError; use ironclaw::llm::{ CompletionRequest, CompletionResponse, FinishReason, LlmProvider, ToolCompletionRequest, @@ -179,37 +178,11 @@ async fn start_test_server() -> (SocketAddr, Arc, Arc, ) -> (SocketAddr, Arc) { - let state = Arc::new(GatewayState { - msg_tx: tokio::sync::RwLock::new(None), - sse: SseManager::new(), - workspace: None, - session_manager: None, - log_broadcaster: None, - log_level_handle: None, - extension_manager: None, - tool_registry: None, - store: None, - job_manager: None, - prompt_queue: None, - scheduler: None, - user_id: "test-user".to_string(), - shutdown_tx: tokio::sync::RwLock::new(None), - ws_tracker: Some(Arc::new(WsConnectionTracker::new())), - llm_provider: Some(llm_provider), - skill_registry: None, - skill_catalog: None, - chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60), - registry_entries: Vec::new(), - cost_guard: None, - startup_time: std::time::Instant::now(), - }); - - let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); - let bound_addr = start_server(addr, state.clone(), AUTH_TOKEN.to_string()) + TestGatewayBuilder::new() + .llm_provider(llm_provider) + .start(AUTH_TOKEN) .await - .expect("Failed to start test server"); - - (bound_addr, state) + .expect("Failed to start test server") } fn client() -> reqwest::Client { @@ -668,35 +641,10 @@ async fn test_models_no_auth() { #[tokio::test] async fn test_no_llm_provider_returns_503() { // Create state WITHOUT llm_provider - let state = Arc::new(GatewayState { - msg_tx: tokio::sync::RwLock::new(None), - sse: SseManager::new(), - workspace: None, - session_manager: None, - log_broadcaster: None, - log_level_handle: None, - extension_manager: None, - tool_registry: None, - store: None, - job_manager: None, - prompt_queue: None, - scheduler: None, - user_id: "test-user".to_string(), - shutdown_tx: tokio::sync::RwLock::new(None), - ws_tracker: Some(Arc::new(WsConnectionTracker::new())), - llm_provider: None, // No LLM! - skill_registry: None, - skill_catalog: None, - chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60), - registry_entries: Vec::new(), - cost_guard: None, - startup_time: std::time::Instant::now(), - }); - - let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); - let bound_addr = start_server(addr, state, AUTH_TOKEN.to_string()) + let (bound_addr, _state) = TestGatewayBuilder::new() + .start(AUTH_TOKEN) .await - .unwrap(); + .expect("Failed to start test server"); let url = format!("http://{}/v1/chat/completions", bound_addr); let resp = client() diff --git a/tests/workspace_integration.rs b/tests/workspace_integration.rs index dddd95e9..54882ec1 100644 --- a/tests/workspace_integration.rs +++ b/tests/workspace_integration.rs @@ -1,4 +1,4 @@ -#![cfg(feature = "postgres")] +#![cfg(all(feature = "postgres", feature = "integration"))] //! Integration tests for the workspace module. //! //! Requires a running PostgreSQL with pgvector extension. @@ -21,18 +21,6 @@ fn get_pool() -> deadpool_postgres::Pool { .expect("Failed to create pool") } -/// Try to get a connection, returning None if Postgres is unreachable. -/// Tests call this to skip gracefully in CI where no database is available. -async fn try_connect(pool: &deadpool_postgres::Pool) -> Option<()> { - match pool.get().await { - Ok(_) => Some(()), - Err(e) => { - eprintln!("skipping: database unavailable ({e})"); - None - } - } -} - async fn cleanup_user(pool: &deadpool_postgres::Pool, user_id: &str) { let conn = pool.get().await.expect("Failed to get connection"); conn.execute( @@ -46,9 +34,6 @@ async fn cleanup_user(pool: &deadpool_postgres::Pool, user_id: &str) { #[tokio::test] async fn test_workspace_write_and_read() { let pool = get_pool(); - if try_connect(&pool).await.is_none() { - return; - } let user_id = "test_write_read"; cleanup_user(&pool, user_id).await; @@ -74,9 +59,6 @@ async fn test_workspace_write_and_read() { #[tokio::test] async fn test_workspace_append() { let pool = get_pool(); - if try_connect(&pool).await.is_none() { - return; - } let user_id = "test_append"; cleanup_user(&pool, user_id).await; @@ -104,9 +86,6 @@ async fn test_workspace_append() { #[tokio::test] async fn test_workspace_nested_paths() { let pool = get_pool(); - if try_connect(&pool).await.is_none() { - return; - } let user_id = "test_nested"; cleanup_user(&pool, user_id).await; @@ -152,9 +131,6 @@ async fn test_workspace_nested_paths() { #[tokio::test] async fn test_workspace_delete() { let pool = get_pool(); - if try_connect(&pool).await.is_none() { - return; - } let user_id = "test_delete"; cleanup_user(&pool, user_id).await; @@ -179,9 +155,6 @@ async fn test_workspace_delete() { #[tokio::test] async fn test_workspace_memory_operations() { let pool = get_pool(); - if try_connect(&pool).await.is_none() { - return; - } let user_id = "test_memory_ops"; cleanup_user(&pool, user_id).await; @@ -210,9 +183,6 @@ async fn test_workspace_memory_operations() { #[tokio::test] async fn test_workspace_daily_log() { let pool = get_pool(); - if try_connect(&pool).await.is_none() { - return; - } let user_id = "test_daily_log"; cleanup_user(&pool, user_id).await; @@ -239,9 +209,6 @@ async fn test_workspace_daily_log() { #[tokio::test] async fn test_workspace_fts_search() { let pool = get_pool(); - if try_connect(&pool).await.is_none() { - return; - } let user_id = "test_fts_search"; cleanup_user(&pool, user_id).await; @@ -300,9 +267,6 @@ async fn test_workspace_fts_search() { #[tokio::test] async fn test_workspace_hybrid_search_with_mock_embeddings() { let pool = get_pool(); - if try_connect(&pool).await.is_none() { - return; - } let user_id = "test_hybrid_search"; cleanup_user(&pool, user_id).await; @@ -342,9 +306,6 @@ async fn test_workspace_hybrid_search_with_mock_embeddings() { #[tokio::test] async fn test_workspace_list_all() { let pool = get_pool(); - if try_connect(&pool).await.is_none() { - return; - } let user_id = "test_list_all"; cleanup_user(&pool, user_id).await; @@ -370,9 +331,6 @@ async fn test_workspace_list_all() { #[tokio::test] async fn test_workspace_system_prompt() { let pool = get_pool(); - if try_connect(&pool).await.is_none() { - return; - } let user_id = "test_system_prompt"; cleanup_user(&pool, user_id).await; diff --git a/tests/ws_gateway_integration.rs b/tests/ws_gateway_integration.rs index 0016ba4e..e95f7a23 100644 --- a/tests/ws_gateway_integration.rs +++ b/tests/ws_gateway_integration.rs @@ -20,10 +20,9 @@ use tokio_tungstenite::tungstenite::Message; use tokio_tungstenite::tungstenite::client::IntoClientRequest; use ironclaw::channels::IncomingMessage; -use ironclaw::channels::web::server::{GatewayState, start_server}; -use ironclaw::channels::web::sse::SseManager; +use ironclaw::channels::web::server::GatewayState; +use ironclaw::channels::web::test_helpers::TestGatewayBuilder; use ironclaw::channels::web::types::SseEvent; -use ironclaw::channels::web::ws::WsConnectionTracker; const AUTH_TOKEN: &str = "test-token-12345"; const TIMEOUT: Duration = Duration::from_secs(5); @@ -37,37 +36,13 @@ async fn start_test_server() -> ( ) { let (agent_tx, agent_rx) = mpsc::channel(64); - let state = Arc::new(GatewayState { - msg_tx: tokio::sync::RwLock::new(Some(agent_tx)), - sse: SseManager::new(), - workspace: None, - session_manager: None, - log_broadcaster: None, - log_level_handle: None, - extension_manager: None, - tool_registry: None, - store: None, - job_manager: None, - prompt_queue: None, - scheduler: None, - user_id: "test-user".to_string(), - shutdown_tx: tokio::sync::RwLock::new(None), - ws_tracker: Some(Arc::new(WsConnectionTracker::new())), - llm_provider: None, - skill_registry: None, - skill_catalog: None, - chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60), - registry_entries: Vec::new(), - cost_guard: None, - startup_time: std::time::Instant::now(), - }); - - let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); - let bound_addr = start_server(addr, state.clone(), AUTH_TOKEN.to_string()) + let (addr, state) = TestGatewayBuilder::new() + .msg_tx(agent_tx) + .start(AUTH_TOKEN) .await .expect("Failed to start test server"); - (bound_addr, state, agent_rx) + (addr, state, agent_rx) } /// Connect a WebSocket client with auth token in query parameter.