diff --git a/AGENTS.md b/AGENTS.md index f09d6c3..7955fac 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -58,6 +58,10 @@ is sufficient, whether a timeout applies, whether an external effect needs approval — it belongs to the host, whose threat model and configuration the decision depends on. +`crates/tinytools-std` holds host-independent tool building blocks (cross-agent +file staleness tracking, SSRF-safe URL validation, PATH probing). It is not part of the dependency-light vocabulary crate and may pull +in `tokio`, `parking_lot` and `log`. + Add a crate by creating `crates//` — `members = ["crates/*"]` picks it up by existing. Inherit `version`, `edition`, `rust-version`, `license`, and `repository` from `[workspace.package]`, take shared dependencies from diff --git a/Cargo.lock b/Cargo.lock index 5cd1abc..5af3e5a 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -28,18 +28,74 @@ dependencies = [ "syn", ] +[[package]] +name = "bitflags" +version = "2.13.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3ded4057c258ba199e2d26386d3af3780957ecaee6c4ef4041c6b4b8b97c0b06" + +[[package]] +name = "cfg-if" +version = "1.0.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4e7648175b45a9a48536d676f68d918270699102aa8dab5496df06904c914600" + [[package]] name = "itoa" version = "1.0.18" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" +[[package]] +name = "libc" +version = "0.2.189" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3eaf3ede3fee6db1a4c2ee091bf8a8b4dccdc6d17f656fb07896ee72867612f2" + +[[package]] +name = "lock_api" +version = "0.4.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "224399e74b87b5f3557511d98dff8b14089b3dadafcab6bb93eab67d3aace965" +dependencies = [ + "scopeguard", +] + +[[package]] +name = "log" +version = "0.4.34" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f9f8bd3e56ce4dfc153cf470fffbfa98c7620958b312ca5c3a4b8d5181fd13c6" + [[package]] name = "memchr" version = "2.8.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cf8baf1c55e62ffcace7a9f06f4bd9cd3f0c4beb022d3b367256b91b87513d98" +[[package]] +name = "parking_lot" +version = "0.12.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "93857453250e3077bd71ff98b6a65ea6621a19bb0f559a85248955ac12c45a1a" +dependencies = [ + "lock_api", + "parking_lot_core", +] + +[[package]] +name = "parking_lot_core" +version = "0.9.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2621685985a2ebf1c516881c026032ac7deafcda1a2c9b7850dc81e3dfcb64c1" +dependencies = [ + "cfg-if", + "libc", + "redox_syscall", + "smallvec", + "windows-link", +] + [[package]] name = "pin-project-lite" version = "0.2.17" @@ -64,6 +120,15 @@ dependencies = [ "proc-macro2", ] +[[package]] +name = "redox_syscall" +version = "0.5.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ed2bf2547551a7053d6fdfafda3f938979645c44812fbfcda098faae3f1a362d" +dependencies = [ + "bitflags", +] + [[package]] name = "regex" version = "1.13.1" @@ -93,6 +158,12 @@ version = "0.8.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d6f6ff9a378485b298a5286656da665ba74413d36db0979633275d2e708145d4" +[[package]] +name = "scopeguard" +version = "1.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49" + [[package]] name = "serde" version = "1.0.229" @@ -136,6 +207,12 @@ dependencies = [ "zmij", ] +[[package]] +name = "smallvec" +version = "1.16.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f9395f0f0eee849a9b707b2f06bb92a6a422090e2123bb2ef8e87a0e61892a8e" + [[package]] name = "syn" version = "3.0.4" @@ -179,6 +256,20 @@ dependencies = [ "tracing", ] +[[package]] +name = "tinytools-std" +version = "0.4.1" +dependencies = [ + "anyhow", + "async-trait", + "log", + "parking_lot", + "serde_json", + "tinytools", + "tokio", + "tracing", +] + [[package]] name = "tokio" version = "1.53.1" @@ -222,6 +313,12 @@ version = "1.0.24" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" +[[package]] +name = "windows-link" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5" + [[package]] name = "zmij" version = "1.0.23" diff --git a/crates/tinytools-agent/src/parse/test/mod.rs b/crates/tinytools-agent/src/parse/test/mod.rs index 73eafbd..3705ce3 100644 --- a/crates/tinytools-agent/src/parse/test/mod.rs +++ b/crates/tinytools-agent/src/parse/test/mod.rs @@ -8,6 +8,7 @@ mod function_call; mod glm; mod harmony_mistral; mod invoke_xml; +mod regressions; mod sentinel; mod tagged; diff --git a/crates/tinytools-agent/src/parse/test/regressions.rs b/crates/tinytools-agent/src/parse/test/regressions.rs new file mode 100644 index 0000000..ec3f8a5 --- /dev/null +++ b/crates/tinytools-agent/src/parse/test/regressions.rs @@ -0,0 +1,132 @@ +//! Host-reported regressions: parser behaviors an embedding host pinned with +//! its own tests before the parsers lived here. + +use super::parse; +use crate::types::CallSource; + +#[test] +fn very_large_arguments_still_parse() { + let large_arg = "x".repeat(100_000); + let response = format!( + r#"{{"name":"echo","arguments":{{"message":"{large_arg}"}}}}"# + ); + let (_text, calls) = parse(&response); + assert_eq!(calls.len(), 1, "large arguments should still parse"); + assert_eq!(calls[0].name, "echo"); +} + +#[test] +fn special_characters_in_arguments_survive() { + let response = r#"{"name":"echo","arguments":{"message":"hello \"world\" <>&'\n\t"}}"#; + let (_text, calls) = parse(response); + assert_eq!(calls.len(), 1); + assert_eq!(calls[0].name, "echo"); + assert_eq!(calls[0].arguments["message"], "hello \"world\" <>&'\n\t"); +} + +#[test] +fn cross_alias_closing_tags_are_recovered() { + let response = + "\n{\"name\": \"shell\", \"arguments\": {\"command\": \"date\"}}\n"; + let (text, calls) = parse(response); + assert!(text.is_empty()); + assert_eq!(calls.len(), 1); + assert_eq!(calls[0].name, "shell"); +} + +#[test] +fn raw_tool_json_without_a_wrapper_is_not_a_call() { + // SECURITY: JSON that merely resembles a call, with no wrapper, must not + // execute; otherwise injected content could mimic a tool call. + let response = "Sure, creating the file now.\n{\"name\": \"file_write\", \"arguments\": {\"path\": \"hello.py\", \"content\": \"print('hello')\"}}"; + let (text, calls) = parse(response); + assert!(text.contains("Sure, creating the file now.")); + assert!(calls.is_empty(), "raw JSON without wrappers must not parse"); +} + +#[test] +fn an_empty_tool_result_block_is_not_a_call() { + let response = "I'll run that command.\n\n\n\nDone."; + let (text, calls) = parse(response); + assert!(text.contains("Done.")); + assert!(calls.is_empty()); +} + +#[test] +fn an_empty_tool_calls_array_is_returned_as_text() { + let response = r#"{"content": "Hello", "tool_calls": []}"#; + let (text, calls) = parse(response); + assert!(text.contains("Hello")); + assert!(calls.is_empty()); +} + +#[test] +fn invoke_tag_with_a_json_body_parses_like_tool_call() { + let input = "Some text\n{\"name\":\"echo\",\"arguments\":{\"value\":\"hi\"}}\ntrailing"; + let (text, calls) = parse(input); + assert_eq!(calls.len(), 1); + assert_eq!(calls[0].name, "echo"); + assert_eq!(calls[0].arguments, serde_json::json!({"value": "hi"})); + assert!(text.contains("Some text")); + assert!(text.contains("trailing")); +} + +#[test] +fn invoke_attribute_form_does_not_leak_markup() { + let input = + "Sure.\n\nhi\n\ndone"; + let (text, calls) = parse(input); + assert_eq!(calls.len(), 1); + assert_eq!(calls[0].arguments, serde_json::json!({"value": "hi"})); + assert!(text.contains("Sure.") && text.contains("done")); + assert!(!text.contains("\n", + "rust parsers\n", + "5\n", + "true\n", + "ignored\n", + "" + ); + let (_text, calls) = parse(input); + assert_eq!(calls.len(), 1); + assert_eq!( + calls[0].arguments, + serde_json::json!({"query": "rust parsers", "limit": 5, "fuzzy": true}) + ); +} + +#[test] +fn invoke_without_a_name_attribute_is_not_a_call() { + let input = "\nhi\n"; + let (_text, calls) = parse(input); + assert!(calls.is_empty()); +} + +#[test] +fn tool_call_json_and_invoke_attribute_blocks_mix_in_source_order() { + let input = concat!( + "{\"name\":\"first\",\"arguments\":{\"a\":1}}\n", + "\ntwo\n" + ); + let (_text, calls) = parse(input); + assert_eq!(calls.len(), 2); + assert_eq!(calls[0].name, "first"); + assert_eq!(calls[0].arguments, serde_json::json!({"a": 1})); + assert_eq!(calls[1].name, "second"); + assert_eq!(calls[1].arguments, serde_json::json!({"b": "two"})); +} + +#[test] +fn markdown_fence_with_a_json_body_parses() { + let input = "preamble\n```tool_call\n{\"name\":\"ping\",\"arguments\":{}}\n```\npostamble"; + let (text, calls) = parse(input); + assert_eq!(calls.len(), 1); + assert_eq!(calls[0].name, "ping"); + assert!(calls[0].source != CallSource::Native); + assert!(text.contains("preamble") && text.contains("postamble")); +} diff --git a/crates/tinytools-agent/src/pformat_test.rs b/crates/tinytools-agent/src/pformat_test.rs index c498e1e..4dcb8db 100644 --- a/crates/tinytools-agent/src/pformat_test.rs +++ b/crates/tinytools-agent/src/pformat_test.rs @@ -312,3 +312,43 @@ fn signature_round_trips_with_parser() { assert_eq!(args["location"], json!("Berlin")); assert_eq!(args["unit"], json!("imperial")); } + +fn echo_schema() -> serde_json::Value { + json!({ + "type": "object", + "properties": { + "value": { "type": "string" }, + "count": { "type": "integer" } + } + }) +} + +#[test] +fn build_registry_keys_on_the_tools_own_names() { + let reg = build_registry([("echo", echo_schema()), ("shell", echo_schema())]); + assert!(reg.contains_key("echo")); + assert!(reg.contains_key("shell")); + assert_eq!(reg.len(), 2); +} + +#[test] +fn a_tool_absent_from_the_registry_cannot_be_called_by_guessing_its_name() { + // The parser must not invent argument names for a tool it does not know, + // or a model could tunnel arbitrary JSON through by guessing a name. + let reg = build_registry([("echo", echo_schema())]); + assert!(parse_call("shell[rm -rf /]", ®).is_none()); +} + +#[test] +fn a_built_registry_parses_positionally_with_schema_ordered_slots() { + let reg = build_registry([("echo", echo_schema())]); + let (name, args) = parse_call("echo[0|3|1|hi]", ®).expect("known tool parses"); + assert_eq!(name, "echo"); + // Schema properties are ordered alphabetically: count, value. + assert_eq!(args["count"], 3); + assert_eq!(args["value"], "hi"); + assert_eq!( + render_signature_from_schema("echo", &echo_schema()), + "echo[0||1|]" + ); +} diff --git a/crates/tinytools-std/Cargo.toml b/crates/tinytools-std/Cargo.toml new file mode 100644 index 0000000..c48e0cb --- /dev/null +++ b/crates/tinytools-std/Cargo.toml @@ -0,0 +1,27 @@ +[package] +name = "tinytools-std" +version.workspace = true +edition.workspace = true +rust-version.workspace = true +license.workspace = true +repository.workspace = true +description = "Host-independent building blocks for agent tools: cross-agent file staleness tracking, SSRF-safe URL validation, and PATH probing." +documentation = "https://docs.rs/tinytools-std" +readme = "README.md" +publish = false + +[dependencies] +anyhow = { workspace = true } +async-trait = { workspace = true } +log = "0.4" +parking_lot = "0.12" +serde_json = { workspace = true } +tinytools = { path = "../tinytools", version = "0.4.1" } +tokio = { version = "1", default-features = false, features = ["rt", "sync"] } +tracing = { workspace = true } + +[dev-dependencies] +tokio = { workspace = true } + +[lints] +workspace = true diff --git a/crates/tinytools-std/README.md b/crates/tinytools-std/README.md new file mode 100644 index 0000000..96ddf51 --- /dev/null +++ b/crates/tinytools-std/README.md @@ -0,0 +1,11 @@ +# tinytools-std + +Host-independent building blocks for agent tools, extracted from OpenHuman. + +| Module | What it is | +| --- | --- | +| `file_state` | Process-wide read/write stamps so parallel agents detect stale or partial reads before overwriting a file. The host decides whether the guard is on (`init_global(enabled)`). | +| `url_guard` | URL validation with SSRF and DNS-rebinding checks for outbound network tools. | +| `detect_tools` | `find_on_path` and the read-only `detect_tools` tool. | + +No enforcement of host policy lives here; the crate only supplies mechanisms. diff --git a/crates/tinytools-std/src/detect_tools/mod.rs b/crates/tinytools-std/src/detect_tools/mod.rs new file mode 100644 index 0000000..bf6d25e --- /dev/null +++ b/crates/tinytools-std/src/detect_tools/mod.rs @@ -0,0 +1,152 @@ +//! Tool: `detect_tools` — report which developer toolchains are installed on PATH. +//! +//! Lets the agent ground its plans in what the host actually has rather than +//! assuming. Read-only: it only scans `$PATH` for executables (no subprocesses, +//! no writes), so it is safe in every access mode. + +use async_trait::async_trait; +use serde_json::json; +use std::path::PathBuf; +use tinytools::{PermissionLevel, Tool, ToolResult}; + +/// Common developer tools probed when the caller doesn't specify a list. +const DEFAULT_CANDIDATES: &[&str] = &[ + "node", "npm", "npx", "pnpm", "yarn", "bun", "deno", "python3", "python", "pip3", "pip", "uv", + "pipx", "cargo", "rustc", "go", "gcc", "cc", "clang", "make", "git", "gh", "docker", "podman", + "kubectl", "rg", "jq", "fd", "curl", "wget", +]; + +/// Read-only tool reporting which developer toolchains are on `PATH`. +#[derive(Debug)] +pub struct DetectToolsTool; + +impl DetectToolsTool { + #[must_use] + /// Create the tool. + pub fn new() -> Self { + Self + } +} + +impl Default for DetectToolsTool { + fn default() -> Self { + Self::new() + } +} + +/// Locate `name` on `$PATH`, honoring `PATHEXT` on Windows. Returns the first +/// matching executable path, or `None` if not found. +#[must_use] +pub fn find_on_path(name: &str) -> Option { + let path = std::env::var_os("PATH")?; + let exts: Vec = if cfg!(windows) { + std::env::var("PATHEXT") + .unwrap_or_else(|_| ".EXE;.CMD;.BAT".to_string()) + .split(';') + .map(std::string::ToString::to_string) + .collect() + } else { + vec![String::new()] + }; + for dir in std::env::split_paths(&path) { + for ext in &exts { + let candidate = dir.join(format!("{name}{ext}")); + if candidate.is_file() { + // On Unix a plain `is_file()` can match a non-executable file and + // falsely report the tool as available; require the exec bit. + #[cfg(unix)] + { + use std::os::unix::fs::PermissionsExt; + let is_exec = std::fs::metadata(&candidate) + .is_ok_and(|m| m.permissions().mode() & 0o111 != 0); + if is_exec { + return Some(candidate); + } + } + #[cfg(not(unix))] + { + return Some(candidate); + } + } + } + } + None +} + +#[async_trait] +impl Tool for DetectToolsTool { + fn name(&self) -> &'static str { + "detect_tools" + } + + fn description(&self) -> &'static str { + "Detect which developer tools / language runtimes are installed on the host PATH \ + (e.g. node, python3, cargo, docker, git, rg). Use this before assuming a tool \ + exists or before proposing to install one. Read-only — scans PATH only." + } + + fn parameters_schema(&self) -> serde_json::Value { + json!({ + "type": "object", + "properties": { + "tools": { + "type": "array", + "items": { "type": "string" }, + "description": "Optional list of tool names to probe. If omitted, a default \ + catalog of common developer tools is probed." + } + } + }) + } + + fn permission_level(&self) -> PermissionLevel { + PermissionLevel::ReadOnly + } + + async fn execute(&self, args: serde_json::Value) -> anyhow::Result { + let requested: Vec = args + .get("tools") + .and_then(|v| v.as_array()) + .map(|a| { + a.iter() + .filter_map(|x| x.as_str().map(str::to_string)) + .collect() + }) + .unwrap_or_default(); + let candidates: Vec = if requested.is_empty() { + DEFAULT_CANDIDATES + .iter() + .map(|s| (*s).to_string()) + .collect() + } else { + requested + }; + + let mut available = Vec::new(); + let mut missing = Vec::new(); + for name in &candidates { + match find_on_path(name) { + Some(p) => available.push(json!({ + "name": name, + "path": p.to_string_lossy(), + })), + None => missing.push(name.clone()), + } + } + + tracing::debug!( + probed = candidates.len(), + available = available.len(), + "[detect_tools] PATH scan complete" + ); + let payload = json!({ + "available": available, + "missing": missing, + "probed": candidates.len(), + }); + Ok(ToolResult::success(serde_json::to_string_pretty(&payload)?)) + } +} + +#[cfg(test)] +mod test; diff --git a/crates/tinytools-std/src/detect_tools/test.rs b/crates/tinytools-std/src/detect_tools/test.rs new file mode 100644 index 0000000..2b9fc28 --- /dev/null +++ b/crates/tinytools-std/src/detect_tools/test.rs @@ -0,0 +1,40 @@ +#![allow(clippy::expect_used, clippy::panic, clippy::unwrap_used)] + +use super::*; + +#[test] +fn name_and_permission() { + let tool = DetectToolsTool::new(); + assert_eq!(tool.name(), "detect_tools"); + assert_eq!(tool.permission_level(), PermissionLevel::ReadOnly); +} + +#[tokio::test] +async fn missing_tool_reported_missing() { + let tool = DetectToolsTool::new(); + let result = tool + .execute(json!({ "tools": ["definitely_not_a_real_binary_xyz_123"] })) + .await + .unwrap(); + assert!(!result.is_error); + let payload: serde_json::Value = serde_json::from_str(&result.output()).unwrap(); + assert_eq!(payload["probed"], 1); + assert_eq!(payload["available"].as_array().unwrap().len(), 0); + assert_eq!( + payload["missing"].as_array().unwrap()[0], + "definitely_not_a_real_binary_xyz_123" + ); +} + +#[tokio::test] +async fn available_plus_missing_equals_probed() { + let tool = DetectToolsTool::new(); + let result = tool + .execute(json!({ "tools": ["sh", "definitely_not_a_real_binary_xyz_123"] })) + .await + .unwrap(); + let payload: serde_json::Value = serde_json::from_str(&result.output()).unwrap(); + let avail = payload["available"].as_array().unwrap().len(); + let miss = payload["missing"].as_array().unwrap().len(); + assert_eq!(avail + miss, 2); +} diff --git a/crates/tinytools-std/src/file_state/agent_context.rs b/crates/tinytools-std/src/file_state/agent_context.rs new file mode 100644 index 0000000..1624adf --- /dev/null +++ b/crates/tinytools-std/src/file_state/agent_context.rs @@ -0,0 +1,24 @@ +//! Task-local carrier for the currently-executing agent's identity so +//! file tools can attribute reads/writes without widening the `Tool` trait. +//! +//! Set by the agent harness around tool execution; tools read via [`current_file_state_agent_id`]. + +tokio::task_local! { + static FILE_STATE_AGENT_ID: String; +} + +/// Returns the current agent's identity for file-state tracking, if set. +/// +/// Returns `None` outside an agent turn (CLI, JSON-RPC direct, unit tests). +#[must_use] +pub fn current_file_state_agent_id() -> Option { + FILE_STATE_AGENT_ID.try_with(std::clone::Clone::clone).ok() +} + +/// Run `future` with `agent_id` installed as the file-state identity. +pub async fn with_file_state_agent_id(agent_id: String, future: F) -> R +where + F: std::future::Future, +{ + FILE_STATE_AGENT_ID.scope(agent_id, future).await +} diff --git a/crates/tinytools-std/src/file_state/mod.rs b/crates/tinytools-std/src/file_state/mod.rs new file mode 100644 index 0000000..df24ea1 --- /dev/null +++ b/crates/tinytools-std/src/file_state/mod.rs @@ -0,0 +1,26 @@ +//! Process-wide file state coordinator for cross-agent staleness detection. +//! +//! Parallel subagents and worker threads share a workspace. Without +//! coordination one worker can read a file, a sibling can edit it, and +//! the first worker can later write based on stale content. This module +//! tracks per-agent read stamps and per-path write stamps so that write +//! tools can detect the conflict and return a model-facing error +//! requiring the agent to re-read. +//! +//! The guard is opt-in for the process: the host calls [`init_global`] with +//! `enabled` (typically derived from its own configuration or environment). +//! Until it does, or when it passes `false`, every operation is a no-op. + +mod agent_context; +mod ops; +mod types; + +pub use agent_context::{current_file_state_agent_id, with_file_state_agent_id}; +pub use ops::{ + acquire_path_lock, check_partial_read, check_stale_read, init_global, parent_stale_files, + record_read, record_write, try_global, +}; +pub use types::{FileStateCoordinator, ReadStamp}; + +#[cfg(test)] +mod test; diff --git a/crates/tinytools-std/src/file_state/ops.rs b/crates/tinytools-std/src/file_state/ops.rs new file mode 100644 index 0000000..b0d3266 --- /dev/null +++ b/crates/tinytools-std/src/file_state/ops.rs @@ -0,0 +1,170 @@ +//! Operational API for the file state coordinator. + +use std::path::{Path, PathBuf}; +use std::sync::{Arc, OnceLock}; +use std::time::{Instant, SystemTime}; +use tokio::sync::{Mutex, OwnedMutexGuard}; + +use super::types::{FileStateCoordinator, ReadStamp, WriteStamp}; + +// ── Singleton ──────────────────────────────────────────────────────────── + +static GLOBAL: OnceLock> = OnceLock::new(); + +/// Initialise the process-global coordinator when `enabled`. Safe to call +/// multiple times; only the first enabling call wins. Passing `false` leaves the +/// guard off, so [`try_global`] keeps returning `None` and every tracking +/// operation stays a no-op. +pub fn init_global(enabled: bool) { + if !enabled { + tracing::debug!("[file_state] guard disabled by host"); + return; + } + let _ = GLOBAL.set(Arc::new(FileStateCoordinator::new())); + tracing::debug!("[file_state] coordinator initialised"); +} + +/// Returns the global coordinator, or `None` when disabled / not yet initialised. +pub fn try_global() -> Option> { + GLOBAL.get().cloned() +} + +// ── Read tracking ──────────────────────────────────────────────────────── + +/// Record that `agent_id` read `resolved_path` at the given mtime. +pub fn record_read(agent_id: &str, resolved_path: PathBuf, mtime: SystemTime, partial: bool) { + let Some(coord) = try_global() else { return }; + tracing::trace!( + agent = agent_id, + path = %resolved_path.display(), + partial, + "[file_state] record_read" + ); + coord.reads.write().insert( + (agent_id.to_string(), resolved_path), + ReadStamp { + mtime, + timestamp: Instant::now(), + partial, + }, + ); +} + +// ── Write tracking ─────────────────────────────────────────────────────── + +/// Record that `agent_id` wrote `resolved_path`. +pub fn record_write(agent_id: &str, resolved_path: PathBuf) { + let Some(coord) = try_global() else { return }; + tracing::trace!( + agent = agent_id, + path = %resolved_path.display(), + "[file_state] record_write" + ); + let now = Instant::now(); + coord.writes.write().insert( + resolved_path.clone(), + WriteStamp { + writer: agent_id.to_string(), + timestamp: now, + }, + ); + // Also update this agent's own read stamp so its own subsequent + // writes don't trigger self-staleness. + coord.reads.write().insert( + (agent_id.to_string(), resolved_path), + ReadStamp { + mtime: SystemTime::now(), + timestamp: now, + partial: false, + }, + ); +} + +// ── Staleness checks ───────────────────────────────────────────────────── + +/// Check whether `agent_id`'s view of `resolved_path` is stale because +/// another agent wrote to it after this agent's last read. Returns an +/// error message when stale, `None` when safe. +#[must_use] +pub fn check_stale_read(agent_id: &str, resolved_path: &PathBuf) -> Option { + let coord = try_global()?; + let reads = coord.reads.read(); + let writes = coord.writes.read(); + let read_key = (agent_id.to_string(), resolved_path.clone()); + let read_stamp = reads.get(&read_key)?; + let ws = writes.get(resolved_path)?; + if ws.writer != agent_id && ws.timestamp > read_stamp.timestamp { + let display_path = resolved_path.display(); + Some(format!( + "Stale read: file '{display_path}' was modified by agent '{}' after your last read. \ + Re-read the file before editing.", + ws.writer + )) + } else { + None + } +} + +/// Check whether `agent_id`'s last read of `resolved_path` was partial. +/// Returns an error message when partial, `None` when safe. +#[must_use] +pub fn check_partial_read(agent_id: &str, resolved_path: &Path) -> Option { + let coord = try_global()?; + let reads = coord.reads.read(); + let read_key = (agent_id.to_string(), resolved_path.to_path_buf()); + let read_stamp = reads.get(&read_key)?; + if read_stamp.partial { + let display_path = resolved_path.display(); + Some(format!( + "Partial read: your last read of '{display_path}' was partial (paginated). \ + Perform a full read before overwriting." + )) + } else { + None + } +} + +// ── Path locking ───────────────────────────────────────────────────────── + +/// Acquire an async lock on `resolved_path` for a read-modify-write +/// section. Returns an `OwnedMutexGuard` that releases when dropped. +/// Returns `None` when the coordinator is disabled. +pub async fn acquire_path_lock(resolved_path: &Path) -> Option> { + let coord = try_global()?; + let mutex = { + let mut locks = coord.path_locks.write(); + locks + .entry(resolved_path.to_path_buf()) + .or_insert_with(|| Arc::new(Mutex::new(()))) + .clone() + }; + Some(mutex.lock_owned().await) +} + +// ── Parent reminder ────────────────────────────────────────────────────── + +/// Return resolved paths that `parent_agent_id` had previously read but +/// were subsequently written by any agent in `child_agent_ids`. +#[must_use] +pub fn parent_stale_files(parent_agent_id: &str, child_agent_ids: &[String]) -> Vec { + let Some(coord) = try_global() else { + return Vec::new(); + }; + let reads = coord.reads.read(); + let writes = coord.writes.read(); + let mut stale = Vec::new(); + for ((agent_id, path), read_stamp) in reads.iter() { + if agent_id != parent_agent_id { + continue; + } + if let Some(ws) = writes.get(path) + && child_agent_ids.contains(&ws.writer) + && ws.timestamp > read_stamp.timestamp + { + stale.push(path.clone()); + } + } + stale.sort(); + stale.dedup(); + stale +} diff --git a/crates/tinytools-std/src/file_state/test/agent_context.rs b/crates/tinytools-std/src/file_state/test/agent_context.rs new file mode 100644 index 0000000..4bf301e --- /dev/null +++ b/crates/tinytools-std/src/file_state/test/agent_context.rs @@ -0,0 +1,35 @@ +use crate::file_state::{current_file_state_agent_id, with_file_state_agent_id}; + +#[tokio::test] +async fn returns_none_outside_scope() { + assert_eq!(current_file_state_agent_id(), None); +} + +#[tokio::test] +async fn installs_and_reads_agent_id() { + let observed = + with_file_state_agent_id("agent-1".into(), async { current_file_state_agent_id() }).await; + assert_eq!(observed, Some("agent-1".to_string())); +} + +#[tokio::test] +async fn does_not_leak_across_scopes() { + with_file_state_agent_id("agent-1".into(), async { + assert_eq!(current_file_state_agent_id(), Some("agent-1".to_string())); + }) + .await; + assert_eq!(current_file_state_agent_id(), None); +} + +#[tokio::test] +async fn nested_scope_overrides_outer() { + with_file_state_agent_id("parent".into(), async { + assert_eq!(current_file_state_agent_id(), Some("parent".to_string())); + with_file_state_agent_id("child".into(), async { + assert_eq!(current_file_state_agent_id(), Some("child".to_string())); + }) + .await; + assert_eq!(current_file_state_agent_id(), Some("parent".to_string())); + }) + .await; +} diff --git a/crates/tinytools-std/src/file_state/test/mod.rs b/crates/tinytools-std/src/file_state/test/mod.rs new file mode 100644 index 0000000..4621a9b --- /dev/null +++ b/crates/tinytools-std/src/file_state/test/mod.rs @@ -0,0 +1,4 @@ +#![allow(clippy::expect_used, clippy::panic, clippy::unwrap_used)] + +mod agent_context; +mod ops; diff --git a/crates/tinytools-std/src/file_state/test/ops.rs b/crates/tinytools-std/src/file_state/test/ops.rs new file mode 100644 index 0000000..f879afe --- /dev/null +++ b/crates/tinytools-std/src/file_state/test/ops.rs @@ -0,0 +1,196 @@ +use crate::file_state::types::WriteStamp; +use crate::file_state::{FileStateCoordinator, ReadStamp}; +use std::path::PathBuf; +use std::sync::Arc; +use std::time::{Duration, Instant, SystemTime}; +use tokio::sync::Mutex; + +fn fresh_coordinator() -> Arc { + Arc::new(FileStateCoordinator::new()) +} + +#[test] +fn record_and_check_no_staleness() { + let coord = fresh_coordinator(); + let path = PathBuf::from("/tmp/test/a.txt"); + coord.reads.write().insert( + ("agent-a".to_string(), path.clone()), + ReadStamp { + mtime: SystemTime::now(), + timestamp: Instant::now(), + partial: false, + }, + ); + let reads = coord.reads.read(); + let rs = reads.get(&("agent-a".to_string(), path.clone())).unwrap(); + assert!(!rs.partial); + assert!(coord.writes.read().get(&path).is_none()); +} + +#[test] +fn detect_sibling_write_staleness() { + let coord = fresh_coordinator(); + let path = PathBuf::from("/tmp/test/b.txt"); + let read_time = Instant::now(); + coord.reads.write().insert( + ("agent-a".to_string(), path.clone()), + ReadStamp { + mtime: SystemTime::now(), + timestamp: read_time, + partial: false, + }, + ); + std::thread::sleep(Duration::from_millis(5)); + coord.writes.write().insert( + path.clone(), + WriteStamp { + writer: "agent-b".to_string(), + timestamp: Instant::now(), + }, + ); + let stale = coord.stale_reads_for_parent("agent-a"); + assert_eq!(stale, vec![path]); +} + +#[test] +fn own_write_does_not_trigger_staleness() { + let coord = fresh_coordinator(); + let path = PathBuf::from("/tmp/test/c.txt"); + let now = Instant::now(); + coord.reads.write().insert( + ("agent-a".to_string(), path.clone()), + ReadStamp { + mtime: SystemTime::now(), + timestamp: now, + partial: false, + }, + ); + std::thread::sleep(Duration::from_millis(5)); + coord.writes.write().insert( + path.clone(), + WriteStamp { + writer: "agent-a".to_string(), + timestamp: Instant::now(), + }, + ); + let stale = coord.stale_reads_for_parent("agent-a"); + assert!(stale.is_empty()); +} + +#[test] +fn partial_read_detected() { + let coord = fresh_coordinator(); + let path = PathBuf::from("/tmp/test/d.txt"); + coord.reads.write().insert( + ("agent-a".to_string(), path.clone()), + ReadStamp { + mtime: SystemTime::now(), + timestamp: Instant::now(), + partial: true, + }, + ); + let reads = coord.reads.read(); + let rs = reads.get(&("agent-a".to_string(), path.clone())).unwrap(); + assert!(rs.partial); +} + +#[test] +fn parent_stale_files_detects_child_writes() { + let coord = fresh_coordinator(); + let path = PathBuf::from("/tmp/test/e.txt"); + let parent_read_time = Instant::now(); + coord.reads.write().insert( + ("parent".to_string(), path.clone()), + ReadStamp { + mtime: SystemTime::now(), + timestamp: parent_read_time, + partial: false, + }, + ); + std::thread::sleep(Duration::from_millis(5)); + coord.writes.write().insert( + path.clone(), + WriteStamp { + writer: "child-1".to_string(), + timestamp: Instant::now(), + }, + ); + let stale = coord.stale_reads_for_parent("parent"); + assert_eq!(stale, vec![path]); +} + +#[test] +fn paths_written_by_collects_correctly() { + let coord = fresh_coordinator(); + let p1 = PathBuf::from("/tmp/test/f1.txt"); + let p2 = PathBuf::from("/tmp/test/f2.txt"); + coord.writes.write().insert( + p1.clone(), + WriteStamp { + writer: "child-1".to_string(), + timestamp: Instant::now(), + }, + ); + coord.writes.write().insert( + p2.clone(), + WriteStamp { + writer: "child-2".to_string(), + timestamp: Instant::now(), + }, + ); + let result = coord.paths_written_by(&["child-1".to_string()]); + assert_eq!(result.len(), 1); + assert!(result.contains_key("child-1")); + assert_eq!(result["child-1"], vec![p1]); +} + +#[tokio::test] +async fn path_lock_serialises_access() { + let coord = fresh_coordinator(); + let path = PathBuf::from("/tmp/test/lock.txt"); + let mutex = { + let mut locks = coord.path_locks.write(); + locks + .entry(path.clone()) + .or_insert_with(|| Arc::new(Mutex::new(()))) + .clone() + }; + + let guard = mutex.lock().await; + assert!(mutex.try_lock().is_err()); + drop(guard); + assert!(mutex.try_lock().is_ok()); +} + +#[tokio::test] +async fn global_api_tracks_reads_writes_and_locks() { + use crate::file_state::{ + acquire_path_lock, check_partial_read, check_stale_read, init_global, parent_stale_files, + record_read, record_write, try_global, + }; + + init_global(false); + init_global(true); + assert!(try_global().is_some()); + + let path = PathBuf::from("/tmp/test/global-flow.txt"); + record_read("reader", path.clone(), SystemTime::now(), true); + assert!(check_partial_read("reader", &path).is_some()); + assert!(check_stale_read("reader", &path).is_none()); + + record_read("reader", path.clone(), SystemTime::now(), false); + assert!(check_partial_read("reader", &path).is_none()); + + std::thread::sleep(Duration::from_millis(5)); + record_write("writer", path.clone()); + let msg = check_stale_read("reader", &path).expect("stale after sibling write"); + assert!(msg.contains("writer")); + assert_eq!( + parent_stale_files("reader", &["writer".to_string()]), + vec![path.clone()] + ); + assert!(check_stale_read("writer", &path).is_none()); + + let guard = acquire_path_lock(&path).await; + assert!(guard.is_some()); +} diff --git a/crates/tinytools-std/src/file_state/types.rs b/crates/tinytools-std/src/file_state/types.rs new file mode 100644 index 0000000..56a527e --- /dev/null +++ b/crates/tinytools-std/src/file_state/types.rs @@ -0,0 +1,99 @@ +//! Core types for the file state coordinator. + +use parking_lot::RwLock; +use std::collections::HashMap; +use std::path::PathBuf; +use std::sync::Arc; +use std::time::{Instant, SystemTime}; +use tokio::sync::Mutex; + +/// Snapshot of a single file read by an agent. +#[derive(Debug, Clone)] +pub struct ReadStamp { + /// Filesystem mtime at the moment of the read. + pub mtime: SystemTime, + /// Monotonic clock timestamp of the read. + pub timestamp: Instant, + /// Whether the read was partial (paginated / offset+limit). + pub partial: bool, +} + +/// Per-path write metadata. +#[derive(Debug, Clone)] +pub(crate) struct WriteStamp { + /// Agent identity that performed the write. + pub writer: String, + /// Monotonic clock timestamp of the write. + pub timestamp: Instant, +} + +/// Process-global coordinator that tracks file reads and writes across +/// all agents in the process. Thread-safe via `RwLock`. +#[derive(Debug)] +pub struct FileStateCoordinator { + /// Per-agent, per-resolved-path read stamps. + /// Key: `(agent_id, canonical_path)`. + pub(crate) reads: RwLock>, + + /// Per-resolved-path write stamp (last writer wins). + pub(crate) writes: RwLock>, + + /// Per-resolved-path async mutex for serialising read-modify-write + /// sections (used by `edit` and `apply_patch`). + pub(crate) path_locks: RwLock>>>, +} + +impl Default for FileStateCoordinator { + fn default() -> Self { + Self::new() + } +} + +impl FileStateCoordinator { + #[must_use] + /// Create an empty coordinator. + pub fn new() -> Self { + Self { + reads: RwLock::new(HashMap::new()), + writes: RwLock::new(HashMap::new()), + path_locks: RwLock::new(HashMap::new()), + } + } + + /// Return the set of resolved paths that `parent_agent_id` has read + /// but were subsequently written by a different agent. + pub fn stale_reads_for_parent(&self, parent_agent_id: &str) -> Vec { + let reads = self.reads.read(); + let writes = self.writes.read(); + let mut stale = Vec::new(); + for ((agent_id, path), read_stamp) in reads.iter() { + if agent_id != parent_agent_id { + continue; + } + if let Some(ws) = writes.get(path) + && ws.writer != parent_agent_id + && ws.timestamp > read_stamp.timestamp + { + stale.push(path.clone()); + } + } + stale.sort(); + stale.dedup(); + stale + } + + /// Collect all paths written by agents in the given set. + pub fn paths_written_by(&self, agent_ids: &[String]) -> HashMap> { + let writes = self.writes.read(); + let mut result: HashMap> = HashMap::new(); + for (path, ws) in writes.iter() { + if agent_ids.contains(&ws.writer) { + result + .entry(ws.writer.clone()) + .or_default() + .push(path.clone()); + } + } + result + } +} diff --git a/crates/tinytools-std/src/lib.rs b/crates/tinytools-std/src/lib.rs new file mode 100644 index 0000000..99d6913 --- /dev/null +++ b/crates/tinytools-std/src/lib.rs @@ -0,0 +1,14 @@ +//! Host-independent building blocks for agent tools. +//! +//! `tinytools` is the vocabulary a tool is written against and deliberately +//! carries no behavior. This crate holds the small, reusable mechanisms that +//! several hosts' tools share, none of which encodes a host's policy: +//! +//! - [`file_state`] — cross-agent read/write stamps and per-path locks, so a +//! sibling agent's edit is noticed before a stale overwrite. +//! - [`url_guard`] — URL validation with SSRF and DNS-rebinding checks. +//! - [`detect_tools`] — `PATH` probing and the read-only `detect_tools` tool. + +pub mod detect_tools; +pub mod file_state; +pub mod url_guard; diff --git a/crates/tinytools-std/src/url_guard/mod.rs b/crates/tinytools-std/src/url_guard/mod.rs new file mode 100644 index 0000000..4dcd7c7 --- /dev/null +++ b/crates/tinytools-std/src/url_guard/mod.rs @@ -0,0 +1,395 @@ +//! Shared URL validation + SSRF guards for outbound network tools. +//! +//! Used by `http_request`, `curl`, and any future tool that takes a +//! user-supplied URL. Two allowlist modes: +//! +//! - **Open allowlist** (`allowed_domains` is empty): any public non-private +//! host is permitted. All SSRF guards still apply (loopback / RFC1918 / +//! link-local / multicast / documentation / shared-address / +//! IPv4-mapped IPv6, `localhost` / `*.localhost` / `*.local`). +//! - **Strict allowlist** (`allowed_domains` is non-empty): only the listed +//! domains and their subdomains are permitted. +//! +//! Both modes enforce: http(s) only, no whitespace, no userinfo, no IPv6 hosts. +//! +//! **Alternate IP notations** (octal, hex, decimal): Rust's `IpAddr::parse` +//! rejects them so they are treated as plain hostnames. In strict-allowlist +//! mode they are rejected by the domain check. In open-allowlist mode they +//! pass `validate_url` but are caught by `validate_url_with_dns_check` +//! because they fail real-world DNS resolution. +//! +//! ## DNS Rebinding Protection +//! +//! Hostname validation alone is insufficient: an attacker can register a +//! domain that alternates DNS responses between a public IP (passing the +//! allowlist) and a private IP (e.g. 127.0.0.1). To close this gap, +//! callers should use [`validate_url_with_dns_check`] which resolves the +//! hostname and re-validates the resolved IPs before the request is made. + +use std::future::Future; +use std::net::{IpAddr, ToSocketAddrs}; + +/// Validate a URL against the allowlist + SSRF rules. Returns the +/// original URL on success. +/// +/// # Errors +/// +/// Fails when the URL is empty, contains whitespace, is not `http(s)`, names a +/// local/private host, or (in strict mode) is outside the allowlist. +pub fn validate_url(raw_url: &str, allowed_domains: &[String]) -> anyhow::Result { + let url = raw_url.trim(); + + if url.is_empty() { + anyhow::bail!("URL cannot be empty"); + } + + if url.chars().any(char::is_whitespace) { + anyhow::bail!("URL cannot contain whitespace"); + } + + if !url.starts_with("http://") && !url.starts_with("https://") { + anyhow::bail!("Only http:// and https:// URLs are allowed"); + } + + let host = extract_host(url)?; + + if is_private_or_local_host(&host) { + log::debug!( + "[url_guard] ssrf block: host={host} mode={}", + if allowed_domains.is_empty() { + "open" + } else { + "strict" + } + ); + anyhow::bail!("Blocked local/private host: {host}"); + } + + // Empty allowed_domains = open mode: any public non-private host is + // permitted (same as ["*"]). This ensures the http_request tool works + // out of the box regardless of whether the user configured an explicit + // domain list, and keeps web-fetch consistent across routing paths. + // A non-empty list = strict mode: only listed domains pass. (#2700) + if !allowed_domains.is_empty() && !host_matches_allowlist(&host, allowed_domains) { + log::debug!( + "[url_guard] strict-allowlist rejection: host={host} allowed={allowed_domains:?}" + ); + anyhow::bail!( + "I'm not allowed to open '{host}' — it isn't in your allowed websites. \ + Add it (or turn on \"Allow all sites\") under \ + Settings → Advanced → Search engine → Allowed websites, then ask me again." + ); + } + + log::debug!( + "[url_guard] validate_url ok: host={host} mode={}", + if allowed_domains.is_empty() { + "open" + } else { + "strict" + } + ); + + Ok(url.to_string()) +} + +/// Like [`validate_url`] but also resolves the hostname via DNS and +/// verifies that none of the resolved IPs are private/local. This +/// defends against DNS rebinding attacks where an attacker's domain +/// initially resolves to a public IP (passing the allowlist) and then +/// flips to 127.0.0.1 at request time. +/// +/// Callers should use this function instead of `validate_url` in all +/// paths that make outbound HTTP requests. +/// +/// # Errors +/// +/// Everything [`validate_url`] rejects, plus DNS failure, an empty answer, or +/// any resolved address that is private/local. +pub async fn validate_url_with_dns_check( + raw_url: &str, + allowed_domains: &[String], +) -> anyhow::Result { + validate_url_with_dns_check_with_resolver(raw_url, allowed_domains, resolve_host_ips).await +} + +async fn validate_url_with_dns_check_with_resolver( + raw_url: &str, + allowed_domains: &[String], + resolver: F, +) -> anyhow::Result +where + F: FnOnce(String, u16) -> Fut, + Fut: Future>>, +{ + let url = validate_url(raw_url, allowed_domains)?; + + let host = extract_host(&url)?; + + // If the host is already a valid IP literal, `is_private_or_local_host` + // has already checked it above. We only need DNS resolution for hostnames. + if host.parse::().is_ok() { + return Ok(url); + } + + let port = extract_port(&url)?; + log::debug!("[url_guard] resolving DNS for host={host} port={port}"); + let addrs = resolver(host.clone(), port).await?; + + if addrs.is_empty() { + anyhow::bail!("DNS resolution returned no addresses for '{host}'"); + } + + log::debug!("[url_guard] DNS resolved host={host} addrs={}", addrs.len()); + + for addr in &addrs { + let ip_str = addr.to_string(); + if is_private_or_local_host(&ip_str) { + log::debug!("[url_guard] DNS rebinding blocked host={host} resolved_ip={ip_str}"); + anyhow::bail!( + "DNS rebinding blocked: '{host}' resolved to private/local address {ip_str}" + ); + } + } + + Ok(url) +} + +async fn resolve_host_ips(host: String, port: u16) -> anyhow::Result> { + let log_host = host.clone(); + tokio::task::spawn_blocking(move || { + (host.as_str(), port) + .to_socket_addrs() + .map_err(|e| { + log::debug!("[url_guard] DNS resolution failed host={host} port={port} error={e}"); + anyhow::anyhow!("DNS resolution failed for '{host}': {e}") + }) + .map(|iter| iter.map(|addr| addr.ip()).collect()) + }) + .await + .map_err(|e| { + log::debug!("[url_guard] DNS resolution task failed host={log_host} port={port} error={e}"); + anyhow::anyhow!("DNS resolution task failed for '{log_host}': {e}") + })? +} + +#[must_use] +/// Normalise an allowlist: strip scheme/path, lowercase, drop invalid entries and +/// duplicates. An empty result means open mode. +pub fn normalize_allowed_domains(domains: Vec) -> Vec { + if domains.is_empty() { + return Vec::new(); + } + let mut normalized = domains + .into_iter() + .filter_map(|d| normalize_domain(&d)) + .collect::>(); + normalized.sort_unstable(); + normalized.dedup(); + if normalized.is_empty() { + // All entries were malformed (whitespace-only, scheme-only, etc.) and + // filtered out. Returning empty would silently enter open mode; instead + // return a sentinel that keeps the tool in strict mode and rejects every + // URL — fail-closed on misconfiguration. (#2738) + log::warn!( + "[url_guard] all configured allowed_domains entries are invalid — \ + treating as misconfigured allowlist (fail-closed)" + ); + return vec!["".to_string()]; + } + normalized +} + +/// Normalise one allowlist entry to a bare lowercase host, or `None` if invalid. +pub fn normalize_domain(raw: &str) -> Option { + let mut d = raw.trim().to_lowercase(); + if d.is_empty() { + return None; + } + + if let Some(stripped) = d.strip_prefix("https://") { + d = stripped.to_string(); + } else if let Some(stripped) = d.strip_prefix("http://") { + d = stripped.to_string(); + } + + if let Some((host, _)) = d.split_once('/') { + d = host.to_string(); + } + + d = d.trim_start_matches('.').trim_end_matches('.').to_string(); + + if let Some((host, _)) = d.split_once(':') { + d = host.to_string(); + } + + if d.is_empty() || d.chars().any(char::is_whitespace) { + return None; + } + + Some(d) +} + +/// Extract the host part of an `http(s)` URL. +/// +/// # Errors +/// +/// Fails on a missing/empty host, userinfo, or an IPv6 literal. +pub fn extract_host(url: &str) -> anyhow::Result { + let rest = url + .strip_prefix("http://") + .or_else(|| url.strip_prefix("https://")) + .ok_or_else(|| anyhow::anyhow!("Only http:// and https:// URLs are allowed"))?; + + let authority = rest + .split(['/', '?', '#']) + .next() + .ok_or_else(|| anyhow::anyhow!("Invalid URL"))?; + + if authority.is_empty() { + anyhow::bail!("URL must include a host"); + } + + if authority.contains('@') { + anyhow::bail!("URL userinfo is not allowed"); + } + + if authority.starts_with('[') { + anyhow::bail!("IPv6 hosts are not supported in http_request"); + } + + let host = authority + .split(':') + .next() + .unwrap_or_default() + .trim() + .trim_end_matches('.') + .to_lowercase(); + + if host.is_empty() { + anyhow::bail!("URL must include a valid host"); + } + + Ok(host) +} + +/// Extract the explicit or scheme-default port of an `http(s)` URL. +/// +/// # Errors +/// +/// Fails when the URL has no valid port. +pub fn extract_port(url: &str) -> anyhow::Result { + let is_http = url.starts_with("http://"); + let rest = url + .strip_prefix("http://") + .or_else(|| url.strip_prefix("https://")) + .ok_or_else(|| anyhow::anyhow!("Only http:// and https:// URLs are allowed"))?; + + let authority = rest + .split(['/', '?', '#']) + .next() + .ok_or_else(|| anyhow::anyhow!("Invalid URL"))?; + + if authority.starts_with('[') { + anyhow::bail!("IPv6 hosts are not supported in http_request"); + } + + if let Some((_, port)) = authority.rsplit_once(':') { + if port.is_empty() || !port.chars().all(|ch| ch.is_ascii_digit()) { + anyhow::bail!("URL port must be numeric"); + } + return port + .parse::() + .map_err(|_| anyhow::anyhow!("URL port is out of range")); + } + + Ok(if is_http { 80 } else { 443 }) +} + +#[must_use] +/// Whether `host` equals, or is a subdomain of, an allowlist entry. +pub fn host_matches_allowlist(host: &str, allowed_domains: &[String]) -> bool { + allowed_domains.iter().any(|domain| { + // `"*"` is the explicit allow-all wildcard (the "Allow all sites" + // toggle), mirroring the browser tool. Local/private hosts are still + // rejected upstream by `is_private_or_local_host`, so a wildcard only + // opens *public* hosts, never the loopback/RFC1918 SSRF surface. + domain == "*" + || host == domain + || host + .strip_suffix(domain) + .is_some_and(|prefix| prefix.ends_with('.')) + }) +} + +#[must_use] +/// Whether `host` is a local name or resolves lexically to a non-global address. +pub fn is_private_or_local_host(host: &str) -> bool { + let unbracketed = host + .strip_prefix('[') + .and_then(|h| h.strip_suffix(']')) + .unwrap_or(host); + let bare = unbracketed.strip_suffix('.').unwrap_or(unbracketed); + + let lower = bare.to_ascii_lowercase(); + + let has_local_tld = lower + .rsplit('.') + .next() + .is_some_and(|label| label == "local"); + + if lower == "localhost" || lower.ends_with(".localhost") || has_local_tld { + return true; + } + + if let Ok(ip) = bare.parse::() { + return match ip { + std::net::IpAddr::V4(v4) => is_non_global_v4(v4), + std::net::IpAddr::V6(v6) => is_non_global_v6(v6), + }; + } + + false +} + +#[must_use] +/// Whether an IPv4 address is non-global (loopback, private, link-local, ...). +pub fn is_non_global_v4(v4: std::net::Ipv4Addr) -> bool { + let [a, b, c, _] = v4.octets(); + v4.is_loopback() + || v4.is_private() + || v4.is_link_local() + || v4.is_unspecified() + || v4.is_broadcast() + || v4.is_multicast() + || (a == 100 && (64..=127).contains(&b)) + || a >= 240 + || (a == 192 && b == 0 && c == 0) + || (a == 192 && b == 88 && c == 99) + || (a == 198 && b == 51 && c == 100) + || (a == 203 && b == 0 && c == 113) + || (a == 198 && (18..=19).contains(&b)) + // 0.0.0.0/8 — "this network" (RFC 1122 §3.2.1.3). `is_unspecified()` only + // covers 0.0.0.0 itself, but the whole /8 routes to the local host on + // Linux. Carried over from the `ops_install` copy this replaces. + || a == 0 +} + +/// Whether an IPv6 address is non-global (loopback, ULA, link-local, mapped, ...). +pub fn is_non_global_v6(v6: std::net::Ipv6Addr) -> bool { + let segs = v6.segments(); + v6.is_loopback() + || v6.is_unspecified() + || v6.is_multicast() + || (segs[0] & 0xfe00) == 0xfc00 + || (segs[0] & 0xffc0) == 0xfe80 + || (segs[0] == 0x2001 && segs[1] == 0x0db8) + || (segs[0] == 0x0100 && segs[1] == 0 && segs[2] == 0 && segs[3] <= 1) + || (segs[0] == 0x2001 && segs[1] == 0x0002 && segs[2] == 0) + || (segs[0] & 0xfff0) == 0x3ff0 + || segs[0] == 0x5f00 + || v6.to_ipv4_mapped().is_some_and(is_non_global_v4) +} + +#[cfg(test)] +mod test; diff --git a/crates/tinytools-std/src/url_guard/test.rs b/crates/tinytools-std/src/url_guard/test.rs new file mode 100644 index 0000000..2747038 --- /dev/null +++ b/crates/tinytools-std/src/url_guard/test.rs @@ -0,0 +1,587 @@ +#![allow(clippy::expect_used, clippy::panic, clippy::unwrap_used)] + +use super::*; + +#[test] +fn normalize_domain_strips_scheme_path_and_case() { + let got = normalize_domain(" HTTPS://Docs.Example.com/path ").unwrap(); + assert_eq!(got, "docs.example.com"); +} + +#[test] +fn normalize_allowed_domains_deduplicates() { + let got = normalize_allowed_domains(vec![ + "example.com".into(), + "EXAMPLE.COM".into(), + "https://example.com/".into(), + ]); + assert_eq!(got, vec!["example.com".to_string()]); +} + +#[test] +fn validate_accepts_exact_domain() { + let allow = vec!["example.com".to_string()]; + let got = validate_url("https://example.com/docs", &allow).unwrap(); + assert_eq!(got, "https://example.com/docs"); +} + +#[test] +fn validate_accepts_http() { + let allow = vec!["example.com".to_string()]; + assert!(validate_url("http://example.com", &allow).is_ok()); +} + +#[test] +fn validate_accepts_subdomain() { + let allow = vec!["example.com".to_string()]; + assert!(validate_url("https://api.example.com/v1", &allow).is_ok()); +} + +#[test] +fn validate_rejects_allowlist_miss() { + let allow = vec!["example.com".to_string()]; + let err = validate_url("https://google.com", &allow) + .unwrap_err() + .to_string(); + assert!(err.contains("allowed websites")); +} + +#[test] +fn validate_wildcard_allows_any_public_host() { + let allow = vec!["*".to_string()]; + assert!(validate_url("https://example.com/docs", &allow).is_ok()); + assert!(validate_url("https://www.cnbc.com/markets", &allow).is_ok()); + assert!(validate_url("https://sub.deep.example.org", &allow).is_ok()); +} + +#[test] +fn validate_wildcard_still_blocks_local_and_private() { + // "Allow all sites" must NOT defeat the SSRF guard. + let allow = vec!["*".to_string()]; + assert!( + validate_url("https://localhost:8080", &allow) + .unwrap_err() + .to_string() + .contains("local/private") + ); + assert!( + validate_url("https://192.168.1.5", &allow) + .unwrap_err() + .to_string() + .contains("local/private") + ); +} + +#[test] +fn validate_rejects_localhost() { + let allow = vec!["localhost".to_string()]; + let err = validate_url("https://localhost:8080", &allow) + .unwrap_err() + .to_string(); + assert!(err.contains("local/private")); +} + +#[test] +fn validate_rejects_private_ipv4() { + let allow = vec!["192.168.1.5".to_string()]; + let err = validate_url("https://192.168.1.5", &allow) + .unwrap_err() + .to_string(); + assert!(err.contains("local/private")); +} + +#[test] +fn validate_rejects_whitespace() { + let allow = vec!["example.com".to_string()]; + let err = validate_url("https://example.com/hello world", &allow) + .unwrap_err() + .to_string(); + assert!(err.contains("whitespace")); +} + +#[test] +fn validate_rejects_userinfo() { + let allow = vec!["example.com".to_string()]; + let err = validate_url("https://user@example.com", &allow) + .unwrap_err() + .to_string(); + assert!(err.contains("userinfo")); +} + +// Empty allowed_domains = open mode: any public host is permitted. +// This keeps web-fetch working when no domain list is configured and +// makes behaviour consistent between default and external-LLM routing. +// (#2700) +#[test] +fn validate_empty_allowlist_allows_public_host() { + assert!(validate_url("https://example.com", &[]).is_ok()); + assert!(validate_url("https://www.cnbc.com/markets", &[]).is_ok()); +} + +#[test] +fn validate_empty_allowlist_still_blocks_private_hosts() { + let err = validate_url("https://192.168.1.5", &[]) + .unwrap_err() + .to_string(); + assert!(err.contains("local/private")); + + let err = validate_url("https://localhost", &[]) + .unwrap_err() + .to_string(); + assert!(err.contains("local/private")); +} + +// ── normalize_allowed_domains: fail-closed on malformed-only input ── + +#[test] +fn normalize_all_invalid_entries_stays_fail_closed() { + // A non-empty list that fully normalizes to nothing must NOT produce + // an empty slice (which would silently enter open mode). (#2738) + let got = normalize_allowed_domains(vec![" ".into(), "https://".into()]); + assert!( + !got.is_empty(), + "normalized result must be non-empty to stay in strict mode" + ); + // The sentinel must not match any real public host. + assert!( + !host_matches_allowlist("example.com", &got), + "sentinel must not grant access to real hosts" + ); + assert!( + !host_matches_allowlist("api.example.com", &got), + "sentinel must not grant access to subdomains" + ); +} + +#[test] +fn normalize_empty_input_stays_empty_for_open_mode() { + // Explicitly empty input should return empty (open mode is intentional). + assert!(normalize_allowed_domains(vec![]).is_empty()); +} + +#[tokio::test] +async fn dns_check_with_empty_allowlist_allows_public_resolved_host() { + // Open mode (empty allowlist) must still pass DNS check for public IPs. + let got = validate_url_with_dns_check_with_resolver( + "https://example.com", + &[], + |host, port| async move { + assert_eq!(host, "example.com"); + assert_eq!(port, 443); + Ok(vec!["93.184.216.34".parse().unwrap()]) + }, + ) + .await + .unwrap(); + assert_eq!(got, "https://example.com"); +} + +#[tokio::test] +async fn dns_check_with_empty_allowlist_blocks_private_resolved_ip() { + // Even in open mode, DNS rebinding to a private IP must be blocked. + let err = validate_url_with_dns_check_with_resolver("https://example.com", &[], |_, _| async { + Ok(vec!["10.0.0.1".parse().unwrap()]) + }) + .await + .unwrap_err() + .to_string(); + assert!(err.contains("DNS rebinding blocked")); +} + +#[tokio::test] +async fn dns_check_resolver_failure_is_a_refusal_not_a_pass_through() { + // A resolver error (NXDOMAIN, network down, timeout) must refuse the + // fetch, not fall back to treating the host as unresolved-and-therefore- + // allowed. + let err = validate_url_with_dns_check_with_resolver( + "https://this-host-does-not-exist.invalid", + &[], + |host, _port| async move { anyhow::bail!("DNS resolution failed for '{host}': NXDOMAIN") }, + ) + .await + .unwrap_err() + .to_string(); + assert!(err.contains("DNS resolution failed")); +} + +#[tokio::test] +async fn dns_check_resolver_returning_no_addresses_is_a_refusal() { + // A resolver that answers with zero addresses (some stub resolvers do + // this instead of erroring) must not be treated as "no IPs to check, + // therefore allowed". + let err = validate_url_with_dns_check_with_resolver("https://example.com", &[], |_, _| async { + Ok(Vec::new()) + }) + .await + .unwrap_err() + .to_string(); + assert!(err.contains("DNS resolution returned no addresses")); +} + +#[test] +fn validate_rejects_ftp_scheme() { + let allow = vec!["example.com".to_string()]; + let err = validate_url("ftp://example.com", &allow) + .unwrap_err() + .to_string(); + assert!(err.contains("http://") || err.contains("https://")); +} + +#[test] +fn validate_rejects_empty_url() { + let allow = vec!["example.com".to_string()]; + let err = validate_url("", &allow).unwrap_err().to_string(); + assert!(err.contains("empty")); +} + +#[test] +fn validate_rejects_ipv6_host() { + let allow = vec!["example.com".to_string()]; + let err = validate_url("http://[::1]:8080/path", &allow) + .unwrap_err() + .to_string(); + assert!(err.contains("IPv6")); +} + +#[test] +fn blocks_multicast_ipv4() { + assert!(is_private_or_local_host("224.0.0.1")); + assert!(is_private_or_local_host("239.255.255.255")); +} + +#[test] +fn blocks_broadcast() { + assert!(is_private_or_local_host("255.255.255.255")); +} + +#[test] +fn blocks_reserved_ipv4() { + assert!(is_private_or_local_host("240.0.0.1")); + assert!(is_private_or_local_host("250.1.2.3")); +} + +#[test] +fn blocks_documentation_ranges() { + // TEST-NET-1 is globally routable in this policy; only TEST-NET-2 and + // TEST-NET-3 are classified as non-global here. + assert!(!is_private_or_local_host("192.0.2.1")); + assert!(is_private_or_local_host("198.51.100.1")); + assert!(is_private_or_local_host("203.0.113.1")); +} + +#[test] +fn blocks_benchmarking_range() { + assert!(is_private_or_local_host("198.18.0.1")); + assert!(is_private_or_local_host("198.19.255.255")); +} + +#[test] +fn blocks_ipv6_localhost() { + assert!(is_private_or_local_host("::1")); + assert!(is_private_or_local_host("[::1]")); +} + +#[test] +fn blocks_ipv6_multicast() { + assert!(is_private_or_local_host("ff02::1")); +} + +#[test] +fn blocks_ipv6_link_local() { + assert!(is_private_or_local_host("fe80::1")); +} + +#[test] +fn blocks_ipv6_unique_local() { + assert!(is_private_or_local_host("fd00::1")); +} + +#[test] +fn blocks_ipv4_mapped_ipv6() { + assert!(is_private_or_local_host("::ffff:127.0.0.1")); + assert!(is_private_or_local_host("::ffff:192.168.1.1")); + assert!(is_private_or_local_host("::ffff:10.0.0.1")); +} + +#[test] +fn allows_public_ipv4() { + assert!(!is_private_or_local_host("8.8.8.8")); + assert!(!is_private_or_local_host("1.1.1.1")); + assert!(!is_private_or_local_host("93.184.216.34")); +} + +#[test] +fn blocks_ipv6_documentation_range() { + assert!(is_private_or_local_host("2001:db8::1")); +} + +#[test] +fn allows_public_ipv6() { + assert!(!is_private_or_local_host("2607:f8b0:4004:800::200e")); +} + +#[test] +fn blocks_shared_address_space() { + assert!(is_private_or_local_host("100.64.0.1")); + assert!(is_private_or_local_host("100.127.255.255")); + assert!(!is_private_or_local_host("100.63.0.1")); + assert!(!is_private_or_local_host("100.128.0.1")); +} + +#[test] +fn ssrf_blocks_loopback_127_range() { + assert!(is_private_or_local_host("127.0.0.1")); + assert!(is_private_or_local_host("127.0.0.2")); + assert!(is_private_or_local_host("127.255.255.255")); +} + +#[test] +fn ssrf_blocks_rfc1918_10_range() { + assert!(is_private_or_local_host("10.0.0.1")); + assert!(is_private_or_local_host("10.255.255.255")); +} + +#[test] +fn ssrf_blocks_rfc1918_172_range() { + assert!(is_private_or_local_host("172.16.0.1")); + assert!(is_private_or_local_host("172.31.255.255")); +} + +#[test] +fn ssrf_blocks_unspecified_address() { + assert!(is_private_or_local_host("0.0.0.0")); +} + +#[test] +fn ssrf_blocks_dot_localhost_subdomain() { + assert!(is_private_or_local_host("evil.localhost")); + assert!(is_private_or_local_host("a.b.localhost")); +} + +#[test] +fn ssrf_blocks_dot_local_tld() { + assert!(is_private_or_local_host("service.local")); +} + +#[test] +fn ssrf_ipv6_unspecified() { + assert!(is_private_or_local_host("::")); +} + +// ── Defense-in-depth: alternate IP notations rejected by allowlist +// +// Rust's IpAddr::parse() rejects octal, hex, decimal, and +// zero-padded notations. They fall through as hostnames and get +// rejected by the allowlist instead. These tests pin that +// behaviour so a parser change can't silently re-open SSRF. + +#[test] +fn ssrf_octal_loopback_not_parsed_as_ip() { + assert!(!is_private_or_local_host("0177.0.0.1")); +} + +#[test] +fn ssrf_hex_loopback_not_parsed_as_ip() { + assert!(!is_private_or_local_host("0x7f000001")); +} + +#[test] +fn ssrf_decimal_loopback_not_parsed_as_ip() { + assert!(!is_private_or_local_host("2130706433")); +} + +#[test] +fn ssrf_zero_padded_loopback_not_parsed_as_ip() { + assert!(!is_private_or_local_host("127.000.000.001")); +} + +#[test] +fn ssrf_alternate_notations_rejected_by_validate_url() { + let allow = vec!["example.com".to_string()]; + for notation in [ + "http://0177.0.0.1", + "http://0x7f000001", + "http://2130706433", + "http://127.000.000.001", + ] { + let err = validate_url(notation, &allow).unwrap_err().to_string(); + assert!( + err.contains("allowed websites"), + "Expected allowlist rejection for {notation}, got: {err}" + ); + } +} + +// ── DNS rebinding protection ───────────────────────────────── + +#[tokio::test] +async fn dns_check_blocks_localhost_resolution() { + // "localhost" resolves to 127.0.0.1 on most systems. Even if + // someone adds it to the allowlist, the DNS check should block it. + let allow = vec!["localhost".to_string()]; + // validate_url itself already blocks "localhost" via the hostname check, + // but validate_url_with_dns_check should also catch it. + let err = validate_url_with_dns_check("https://localhost", &allow) + .await + .unwrap_err() + .to_string(); + assert!( + err.contains("local/private") || err.contains("rebinding"), + "Expected SSRF block for localhost, got: {err}" + ); +} + +#[tokio::test] +async fn dns_check_passes_for_public_resolved_ip() { + let allow = vec!["example.com".to_string()]; + let got = validate_url_with_dns_check_with_resolver( + "https://example.com", + &allow, + |host, port| async move { + assert_eq!(host, "example.com"); + assert_eq!(port, 443); + Ok(vec!["93.184.216.34".parse().unwrap()]) + }, + ) + .await + .unwrap(); + assert_eq!(got, "https://example.com"); +} + +#[tokio::test] +async fn dns_check_blocks_private_resolved_ip() { + let allow = vec!["example.com".to_string()]; + let err = + validate_url_with_dns_check_with_resolver("https://example.com", &allow, |_, _| async { + Ok(vec!["127.0.0.1".parse().unwrap()]) + }) + .await + .unwrap_err() + .to_string(); + assert!(err.contains("DNS rebinding blocked")); +} + +#[tokio::test] +async fn dns_check_uses_explicit_port_for_resolution() { + let allow = vec!["api.example.com".to_string()]; + let got = validate_url_with_dns_check_with_resolver( + "http://api.example.com:8080/status", + &allow, + |host, port| async move { + assert_eq!(host, "api.example.com"); + assert_eq!(port, 8080); + Ok(vec!["93.184.216.34".parse().unwrap()]) + }, + ) + .await + .unwrap(); + assert_eq!(got, "http://api.example.com:8080/status"); +} + +#[tokio::test] +async fn dns_check_returns_resolver_failure() { + let allow = vec!["example.com".to_string()]; + let err = validate_url_with_dns_check_with_resolver( + "https://example.com", + &allow, + |host, _| async move { + anyhow::bail!("DNS resolution failed for '{host}': resolver unavailable") + }, + ) + .await + .unwrap_err() + .to_string(); + assert!(err.contains("DNS resolution failed")); +} + +#[tokio::test] +async fn dns_check_rejects_ip_literal_private() { + let allow = vec!["10.0.0.1".to_string()]; + let err = validate_url_with_dns_check("https://10.0.0.1", &allow) + .await + .unwrap_err() + .to_string(); + assert!(err.contains("local/private")); +} + +#[test] +fn wildcard_allows_any_host() { + let any = vec!["*".to_string()]; + assert!(host_matches_allowlist("docs.rs", &any)); + assert!(host_matches_allowlist("api.github.com", &any)); + assert!(host_matches_allowlist("whatever.example.org", &any)); +} + +#[tokio::test] +async fn wildcard_still_blocks_private_hosts() { + // `*` opens public hosts only — SSRF block on private/local hosts stays. + let any = vec!["*".to_string()]; + let err = validate_url_with_dns_check("https://127.0.0.1", &any) + .await + .unwrap_err() + .to_string(); + assert!(err.contains("local/private"), "got: {err}"); +} + +#[test] +fn exported_ssrf_predicates_classify_non_global_ips_accurately() { + use std::net::{Ipv4Addr, Ipv6Addr}; + + // IPv4 Non-global checks + assert!(is_non_global_v4(Ipv4Addr::LOCALHOST)); + assert!(is_non_global_v4(Ipv4Addr::new(10, 0, 0, 1))); + assert!(is_non_global_v4(Ipv4Addr::new(172, 16, 0, 1))); + assert!(is_non_global_v4(Ipv4Addr::new(192, 168, 1, 1))); + assert!(is_non_global_v4(Ipv4Addr::new(169, 254, 1, 1))); + assert!(is_non_global_v4(Ipv4Addr::new(100, 64, 0, 1))); // CGNAT + assert!(is_non_global_v4(Ipv4Addr::new(240, 0, 0, 1))); // Class E + assert!(!is_non_global_v4(Ipv4Addr::new(192, 0, 2, 1))); // TEST-NET-1 is globally routable in this policy + assert!(is_non_global_v4(Ipv4Addr::new(198, 51, 100, 1))); // TEST-NET-2 + assert!(is_non_global_v4(Ipv4Addr::new(203, 0, 113, 1))); // TEST-NET-3 + assert!(is_non_global_v4(Ipv4Addr::new(192, 88, 99, 1))); // 6to4 anycast + assert!(is_non_global_v4(Ipv4Addr::UNSPECIFIED)); // 0.0.0.0/8 + assert!(is_non_global_v4(Ipv4Addr::new(0, 1, 2, 3))); // 0.0.0.0/8 + + // IPv4 Global public IPs + assert!(!is_non_global_v4(Ipv4Addr::new(8, 8, 8, 8))); + assert!(!is_non_global_v4(Ipv4Addr::new(1, 1, 1, 1))); + assert!(!is_non_global_v4(Ipv4Addr::new(140, 82, 121, 4))); + // Non-TEST-NET IPs in adjacent /24 blocks should not be classified as non-global + assert!(!is_non_global_v4(Ipv4Addr::new(198, 51, 1, 1))); + assert!(!is_non_global_v4(Ipv4Addr::new(203, 0, 1, 1))); + assert!(!is_non_global_v4(Ipv4Addr::new(192, 88, 98, 1))); + + // IPv6 Non-global checks + assert!(is_non_global_v6(Ipv6Addr::LOCALHOST)); + assert!(is_non_global_v6(Ipv6Addr::UNSPECIFIED)); + assert!(is_non_global_v6("fc00::1".parse().unwrap())); + assert!(is_non_global_v6("fe80::1".parse().unwrap())); + assert!(is_non_global_v6("2001:db8::1".parse().unwrap())); + assert!(is_non_global_v6("100::1".parse().unwrap())); + assert!(is_non_global_v6("100:0:0:1::1".parse().unwrap())); + assert!(is_non_global_v6("2001:2::1".parse().unwrap())); + assert!(is_non_global_v6("3fff::1".parse().unwrap())); + assert!(is_non_global_v6("5f00::1".parse().unwrap())); + + // IPv6 Global public IPs + assert!(!is_non_global_v6("2606:4700:4700::1111".parse().unwrap())); + assert!(!is_non_global_v6("101::1".parse().unwrap())); + assert!(!is_non_global_v6("100:0:0:2::1".parse().unwrap())); + assert!(!is_non_global_v6("2001:3::1".parse().unwrap())); + assert!(!is_non_global_v6("4000::1".parse().unwrap())); + assert!(!is_non_global_v6("5f01::1".parse().unwrap())); + + // Host helper checks (including ASCII case-insensitivity and trailing dot) + assert!(is_private_or_local_host("localhost")); + assert!(is_private_or_local_host("LOCALHOST")); + assert!(is_private_or_local_host("localhost.")); + assert!(is_private_or_local_host("my-service.localhost")); + assert!(is_private_or_local_host("MY-SERVICE.LOCALHOST")); + assert!(is_private_or_local_host("device.local")); + assert!(is_private_or_local_host("DEVICE.LOCAL")); + assert!(is_private_or_local_host("device.local.")); + assert!(is_private_or_local_host("127.0.0.1")); + assert!(is_private_or_local_host("[::1]")); + assert!(!is_private_or_local_host("github.com")); + assert!(!is_private_or_local_host("api.openai.com")); +}