use super::{ToolResult, ToolResultDisplay, ToolRuntime, url_fetch};
use crate::{
cancellation::AgentCancellation,
output::{is_credential_like_key, redact_sensitive_text},
};
use chrono::{Duration as ChronoDuration, Utc};
use reqwest::{Url, blocking::Client, redirect::Policy};
use serde::Deserialize;
use serde_json::{Value, json};
use sha2::{Digest, Sha256};
use std::{collections::VecDeque, io::Read, time::Duration};
pub(super) const OUTPUT_MAX_BYTES: usize = 48 * 1024;
const RESPONSE_MAX_BYTES: usize = 2 * 1024 * 1024;
const PAGE_MAX_BYTES: usize = 40_000;
const CACHE_MAX_PAGES: usize = 64;
pub(super) const MAX_DOMAINS: usize = 20;
pub(super) const MAX_DOMAIN_CHARS: usize = 253;
#[derive(Debug, Clone, Copy, Deserialize)]
#[serde(rename_all = "lowercase")]
pub(super) enum RecencyFilter {
Day,
Week,
Month,
Year,
}
impl RecencyFilter {
fn start_date(self) -> String {
let days = match self {
Self::Day => 1,
Self::Week => 7,
Self::Month => 30,
Self::Year => 365,
};
(Utc::now().date_naive() - ChronoDuration::days(days)).to_string()
}
}
#[derive(Debug, Clone, Copy, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub(super) enum Operation {
Search,
Open,
Find,
}
#[derive(Debug, Deserialize)]
#[serde(deny_unknown_fields)]
pub(super) struct WebArgs {
pub operation: Operation,
pub query: Option<String>,
pub url: Option<String>,
pub reference: Option<String>,
pub limit: Option<usize>,
pub offset: Option<usize>,
#[serde(rename = "domainFilter")]
pub domain_filter: Option<Vec<String>>,
#[serde(rename = "recencyFilter")]
pub recency_filter: Option<RecencyFilter>,
}
impl WebArgs {
pub(super) fn validate(mut self) -> anyhow::Result<Self> {
for value in [&mut self.query, &mut self.url, &mut self.reference]
.into_iter()
.flatten()
{
*value = value.trim().to_owned();
anyhow::ensure!(
!value.is_empty(),
"query, url and reference must not be blank"
);
}
anyhow::ensure!(
self.query.as_ref().is_none_or(|q| q.chars().count() <= 512),
"query exceeds 512 characters"
);
anyhow::ensure!(
self.url.as_ref().is_none_or(|u| u.len() <= 2048),
"url exceeds 2048 bytes"
);
anyhow::ensure!(
self.reference.as_ref().is_none_or(|r| r.len() <= 80),
"reference exceeds 80 bytes"
);
anyhow::ensure!(
self.limit.is_none_or(|n| (1..=10).contains(&n)),
"limit must be between 1 and 10"
);
anyhow::ensure!(
self.offset.is_none_or(|n| n <= PAGE_MAX_BYTES),
"offset exceeds 40000"
);
if let Some(domains) = &mut self.domain_filter {
anyhow::ensure!(
domains.len() <= MAX_DOMAINS,
"domainFilter must contain at most {MAX_DOMAINS} domains"
);
for domain in domains {
*domain = domain.trim().to_owned();
validate_domain_filter(domain)?;
}
}
anyhow::ensure!(
self.operation == Operation::Search
|| (self.domain_filter.is_none() && self.recency_filter.is_none()),
"domainFilter and recencyFilter are search-only; omit them for open/find"
);
match self.operation {
Operation::Search => anyhow::ensure!(
self.query.is_some()
&& self.url.is_none()
&& self.reference.is_none()
&& self.offset.is_none(),
"search requires query; url, reference and offset must be omitted"
),
Operation::Open => anyhow::ensure!(
self.query.is_none()
&& (self.url.is_some() != self.reference.is_some())
&& self.limit.is_none()
&& self.offset.is_none(),
"open requires exactly one of url or reference; query, limit and offset must be omitted"
),
Operation::Find => anyhow::ensure!(
self.query.is_some() && self.reference.is_some() && self.url.is_none(),
"find requires query and reference; url must be omitted"
),
}
Ok(self)
}
}
fn validate_domain_filter(domain: &str) -> anyhow::Result<()> {
let domain = domain.strip_prefix('-').unwrap_or(domain);
anyhow::ensure!(
!domain.is_empty() && domain.len() <= MAX_DOMAIN_CHARS,
"domainFilter domain portion must contain 1 to {MAX_DOMAIN_CHARS} characters"
);
for label in domain.split('.') {
anyhow::ensure!(
!label.is_empty()
&& label.len() <= 63
&& !label.starts_with('-')
&& !label.ends_with('-')
&& label
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || byte == b'-'),
"domainFilter entries must be bare ASCII domains with labels of 1 to 63 characters; no schemes, ports, paths, whitespace, or leading/trailing label hyphens"
);
}
Ok(())
}
#[derive(Debug, Clone)]
struct Page {
reference: String,
url: String,
title: String,
text: String,
truncated: bool,
}
#[derive(Debug, Default)]
pub(super) struct WebCache {
pages: VecDeque<Page>,
}
impl WebCache {
fn insert(&mut self, url: String, title: String, text: String) -> Page {
let text = redact_sensitive_text(&text);
let truncated = text.len() > PAGE_MAX_BYTES;
let text = clip(&text, PAGE_MAX_BYTES).to_owned();
let title = clip(&redact_sensitive_text(&title), 512).to_owned();
let mut digest = Sha256::new();
for value in [&url, &title, &text] {
digest.update((value.len() as u64).to_le_bytes());
digest.update(value.as_bytes());
}
digest.update([u8::from(truncated)]);
let digest = digest.finalize();
let reference = format!(
"web:{}",
digest
.iter()
.map(|byte| format!("{byte:02x}"))
.collect::<String>()
);
if let Some(index) = self.pages.iter().position(|p| p.reference == reference) {
self.pages.remove(index);
}
let page = Page {
reference,
url,
title,
text,
truncated,
};
if self.pages.len() == CACHE_MAX_PAGES {
self.pages.pop_front();
}
self.pages.push_back(page.clone());
page
}
fn get(&self, reference: &str) -> anyhow::Result<Page> {
self.pages.iter().find(|p| p.reference == reference).cloned()
.ok_or_else(|| anyhow::anyhow!("web cache miss: reference expired or belongs to another runtime; open the public URL again"))
}
}
#[derive(Deserialize)]
struct Response {
results: Vec<SearchResult>,
}
#[derive(Deserialize)]
struct SearchResult {
url: String,
title: Option<String>,
text: Option<String>,
}
impl ToolRuntime {
pub(super) fn web(
&self,
args: WebArgs,
cancellation: &AgentCancellation,
) -> anyhow::Result<ToolResult> {
let args = args.validate()?;
cancellation.check()?;
if let Some(reference) = &args.reference {
let page = self
.web_cache
.lock()
.map_err(|_| anyhow::anyhow!("web cache lock poisoned"))?
.get(reference)?;
return Ok(if args.operation == Operation::Find {
find_page(
&page,
args.query.as_deref().unwrap_or_default(),
args.limit.unwrap_or(5),
args.offset.unwrap_or(0),
)
} else {
page_result(&page)
});
}
let key = std::env::var("EXA_API_KEY")
.ok()
.filter(|s| !s.trim().is_empty())
.ok_or_else(|| {
anyhow::anyhow!(
"EXA_API_KEY is required for web search/open; set it in the process environment"
)
})?;
if let Some(url) = &args.url {
validate_url(url, cancellation)?;
}
let client = Client::builder()
.connect_timeout(Duration::from_secs(5))
.timeout(Duration::from_secs(20))
.redirect(Policy::none())
.user_agent(concat!("magi-code/", env!("CARGO_PKG_VERSION")))
.build()?;
self.web_with_client(
args,
&client,
"https://api.exa.ai",
key.trim(),
cancellation,
)
}
fn web_with_client(
&self,
args: WebArgs,
client: &Client,
base: &str,
key: &str,
cancellation: &AgentCancellation,
) -> anyhow::Result<ToolResult> {
let (endpoint, body) = match args.operation {
Operation::Search => {
let mut body = json!({"query": args.query, "type": "auto", "numResults": args.limit.unwrap_or(5), "contents": {"text": {"maxCharacters": 10000}}});
let mut include_domains = Vec::new();
let mut exclude_domains = Vec::new();
for domain in args.domain_filter.as_deref().unwrap_or_default() {
if let Some(excluded) = domain.strip_prefix('-') {
exclude_domains.push(excluded);
} else {
include_domains.push(domain.as_str());
}
}
if !include_domains.is_empty() {
body["includeDomains"] = json!(include_domains);
}
if !exclude_domains.is_empty() {
body["excludeDomains"] = json!(exclude_domains);
}
if let Some(recency) = args.recency_filter {
body["startPublishedDate"] = json!(recency.start_date());
}
("search", body)
}
Operation::Open => (
"contents",
json!({"urls": [args.url], "text": {"maxCharacters": 10000}}),
),
Operation::Find => unreachable!("find is cache-only"),
};
cancellation.check()?;
let response = client
.post(format!("{base}/{endpoint}"))
.header("x-api-key", key)
.json(&body)
.send();
cancellation.check()?;
let response = match response {
Ok(response) => response,
Err(error) => {
return Ok(result(
false,
"EXA request failed".into(),
json!({"timeout": error.is_timeout()}),
));
}
};
let status = response.status();
if !status.is_success() {
return Ok(result(
false,
format!("EXA API error {status}"),
json!({"rate_limited": status.as_u16() == 429}),
));
}
anyhow::ensure!(
response
.content_length()
.is_none_or(|n| n <= RESPONSE_MAX_BYTES as u64),
"EXA response exceeds byte limit"
);
let mut bytes = Vec::new();
let read_result = response
.take((RESPONSE_MAX_BYTES + 1) as u64)
.read_to_end(&mut bytes);
cancellation.check()?;
read_result.map_err(|_| anyhow::anyhow!("EXA response body read failed"))?;
anyhow::ensure!(
bytes.len() <= RESPONSE_MAX_BYTES,
"EXA response exceeds byte limit"
);
let response: Response = serde_json::from_slice(&bytes)
.map_err(|_| anyhow::anyhow!("EXA returned malformed JSON"))?;
let mut pages = Vec::new();
let max_results = if args.operation == Operation::Open {
1
} else {
args.limit.unwrap_or(5)
};
let mut omitted = response.results.len().saturating_sub(max_results);
for source in response.results.into_iter().take(max_results) {
cancellation.check()?;
if validate_url(&source.url, cancellation).is_err() {
cancellation.check()?;
omitted += 1;
continue;
}
let Some(text) = source.text.filter(|t| !t.trim().is_empty()) else {
omitted += 1;
continue;
};
let page = self
.web_cache
.lock()
.map_err(|_| anyhow::anyhow!("web cache lock poisoned"))?
.insert(source.url, source.title.unwrap_or_default(), text);
pages.push(page);
}
cancellation.check()?;
if args.operation == Operation::Open {
return Ok(pages.first().map(page_result).unwrap_or_else(|| {
result(
false,
"EXA returned no usable public page content".into(),
json!({"omitted": omitted}),
)
}));
}
let mut content = String::new();
for page in &pages {
content.push_str(&format!(
"[{}] {}\nURL: {}\n{}\n\n",
page.reference,
page.title,
page.url,
clip(&page.text, 1200)
));
}
if pages.is_empty() {
content.push_str("No usable public results returned.");
}
Ok(result(
true,
content,
json!({"results": pages.len(), "omitted": omitted, "truncated": pages.iter().any(|page| page.truncated || page.text.len() > 1200), "scope": "bounded Exa page text; open a reference to inspect cached text"}),
))
}
}
fn validate_url(url: &str, cancellation: &AgentCancellation) -> anyhow::Result<()> {
anyhow::ensure!(url.len() <= 2048, "public URL exceeds 2048 bytes");
let parsed = Url::parse(url).map_err(|_| anyhow::anyhow!("invalid public URL"))?;
anyhow::ensure!(
parsed.username().is_empty() && parsed.password().is_none(),
"URL credentials are not allowed"
);
anyhow::ensure!(
!parsed
.query_pairs()
.any(|(name, _)| is_credential_like_key(&name)),
"sensitive URL parameters are not allowed"
);
if let Some(fragment) = parsed.fragment() {
let mut fragment_url = parsed.clone();
fragment_url.set_query(Some(fragment));
anyhow::ensure!(
!fragment_url
.query_pairs()
.any(|(name, _)| is_credential_like_key(&name)),
"sensitive URL parameters are not allowed"
);
}
anyhow::ensure!(
redact_sensitive_text(url) == url,
"sensitive URL parameters are not allowed"
);
url_fetch::validate_public_url(&parsed, cancellation)
}
fn page_result(page: &Page) -> ToolResult {
result(
true,
format!(
"[{}] {}\nURL: {}\n\n{}",
page.reference, page.title, page.url, page.text
),
json!({"reference": page.reference, "url": page.url, "cache": true, "truncated": page.truncated, "scope": "bounded Exa page text, not necessarily the complete page"}),
)
}
fn find_page(page: &Page, query: &str, limit: usize, offset: usize) -> ToolResult {
let matches: Vec<_> = page
.text
.match_indices(query)
.skip(offset)
.take(limit + 1)
.collect();
let more = matches.len() > limit;
let mut passages = Vec::new();
for (start, matched) in matches.iter().take(limit) {
let mut left = start.saturating_sub(180);
while !page.text.is_char_boundary(left) {
left += 1;
}
let end = (*start + matched.len() + 180).min(page.text.len());
let passage = clip(&page.text[left..], end - left);
passages.push(format!("[{}] byte {}: {}", page.reference, start, passage));
}
let count = passages.len();
let content = if count == 0 {
"No matches in cached text (the cached page may be incomplete).".into()
} else {
passages.join("\n\n")
};
result(
true,
content,
json!({"reference": page.reference, "matches": count, "has_more": more, "next_offset": if more { Some(offset + count) } else { None }, "cache": true, "scope": "case-sensitive literal search of cached text only"}),
)
}
fn clip(text: &str, max: usize) -> &str {
let mut end = text.len().min(max);
while !text.is_char_boundary(end) {
end -= 1;
}
&text[..end]
}
fn result(success: bool, content: String, mut metadata: Value) -> ToolResult {
let content = redact_sensitive_text(&content);
metadata["provider"] = json!("exa");
if content.len() > OUTPUT_MAX_BYTES {
metadata["truncated"] = json!(true);
}
ToolResult {
tool_name: "web".into(),
success,
content: clip(&content, OUTPUT_MAX_BYTES).to_owned(),
metadata,
display: ToolResultDisplay::default(),
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::{io::Write, net::TcpListener, thread};
#[test]
fn search_filters_validate_bounds_and_reject_other_operations_and_aliases() {
let parse = |value| serde_json::from_value::<WebArgs>(value)?.validate();
let longest = format!(
"{}.{}.{}.{}",
"a".repeat(63),
"b".repeat(63),
"c".repeat(63),
"d".repeat(61)
);
let args = parse(json!({"operation":"search","query":"docs","domainFilter":[format!(" {longest} "), format!("-{longest}")]})).unwrap();
assert_eq!(
args.domain_filter.unwrap(),
vec![longest.clone(), format!("-{longest}")]
);
for domain in [
"",
"-",
"--example.com",
"https://example.com",
"example.com/path",
"example.com:443",
"ex ample.com",
".example.com",
"example.com.",
"a..com",
"a-.com",
"é.com",
&"a".repeat(64),
&format!("{longest}e"),
&format!("-{longest}e"),
] {
assert!(
parse(json!({"operation":"search","query":"docs","domainFilter":[domain]}))
.is_err(),
"{domain}"
);
}
assert!(
parse(
json!({"operation":"search","query":"docs","domainFilter":vec!["example.com"; 21]})
)
.is_err()
);
assert!(
parse(
json!({"operation":"search","query":"docs","domainFilter":vec!["example.com"; 20]})
)
.is_ok()
);
for (filter, days) in [("day", 1), ("week", 7), ("month", 30), ("year", 365)] {
let before = Utc::now().date_naive() - ChronoDuration::days(days);
let args =
parse(json!({"operation":"search","query":"docs","recencyFilter":filter})).unwrap();
let actual = args.recency_filter.unwrap().start_date();
let after = Utc::now().date_naive() - ChronoDuration::days(days);
assert!(actual == before.to_string() || actual == after.to_string());
}
for invalid in [
json!({"recencyFilter":"Week"}),
json!({"recencyFilter":"hour"}),
json!({"domain_filter":[]}),
json!({"recency_filter":"week"}),
json!({"queries":["docs"]}),
] {
let mut value = json!({"operation":"search","query":"docs"});
value
.as_object_mut()
.unwrap()
.extend(invalid.as_object().unwrap().clone());
assert!(parse(value).is_err());
}
for base in [
json!({"operation":"open","url":"https://example.com"}),
json!({"operation":"find","reference":"web:test","query":"docs"}),
] {
for (key, value) in [("domainFilter", json!([])), ("recencyFilter", json!("day"))] {
let mut args = base.clone();
args[key] = value;
assert!(parse(args.clone()).is_err());
args[key] = Value::Null;
assert!(parse(args).is_ok());
}
}
}
#[test]
fn exa_timeout_and_cancellation_are_reported_without_body_leaks() {
let temp = tempfile::TempDir::new().unwrap();
let runtime = ToolRuntime::new(temp.path()).unwrap();
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
let server = thread::spawn(move || {
let (mut socket, _) = listener.accept().unwrap();
socket
.set_read_timeout(Some(Duration::from_secs(2)))
.unwrap();
let mut buffer = [0; 4096];
let _ = socket.read(&mut buffer);
thread::sleep(Duration::from_millis(200));
});
let client = Client::builder()
.no_proxy()
.timeout(Duration::from_millis(50))
.build()
.unwrap();
let args = || {
serde_json::from_value::<WebArgs>(json!({"operation":"search", "query":"docs"}))
.unwrap()
.validate()
.unwrap()
};
let result = runtime
.web_with_client(
args(),
&client,
&base,
"test-key",
&AgentCancellation::default(),
)
.unwrap();
server.join().unwrap();
assert!(!result.success);
assert_eq!(result.metadata["timeout"], true);
assert!(!result.content.contains("test-key"));
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
listener.set_nonblocking(true).unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
let (cancellation, handle) = AgentCancellation::default().child_token();
handle.cancel();
let error = runtime
.web_with_client(args(), &client, &base, "test-key", &cancellation)
.unwrap_err();
assert_eq!(error.to_string(), "prompt canceled");
assert_eq!(
listener.accept().unwrap_err().kind(),
std::io::ErrorKind::WouldBlock
);
}
fn mock_call(
operation: &str,
response: String,
) -> (anyhow::Result<ToolResult>, String, ToolRuntime) {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let address = listener.local_addr().unwrap();
let server = thread::spawn(move || {
let (mut socket, _) = listener.accept().unwrap();
socket
.set_read_timeout(Some(Duration::from_secs(2)))
.unwrap();
let mut bytes = Vec::new();
loop {
let mut buffer = [0; 4096];
let count = socket.read(&mut buffer).unwrap();
if count == 0 {
break;
}
bytes.extend_from_slice(&buffer[..count]);
if let Some(end) = bytes.windows(4).position(|w| w == b"\r\n\r\n") {
let header = String::from_utf8_lossy(&bytes[..end]);
let length: usize = header
.lines()
.find_map(|line| {
line.to_ascii_lowercase()
.strip_prefix("content-length:")
.map(|n| n.trim().parse().unwrap())
})
.unwrap_or(0);
if bytes.len() >= end + 4 + length {
break;
}
}
}
let _ = socket.write_all(response.as_bytes());
String::from_utf8(bytes).unwrap()
});
let temp = tempfile::TempDir::new().unwrap();
let runtime = ToolRuntime::new(temp.path()).unwrap();
let value = if operation == "search" {
json!({"operation":"search", "query":"Rust docs", "limit":2, "domainFilter":[" example.com ", "-old.example.com"], "recencyFilter":"week"})
} else {
json!({"operation":"open", "url":"https://1.1.1.1/page"})
};
let args = serde_json::from_value::<WebArgs>(value)
.unwrap()
.validate()
.unwrap();
let client = Client::builder()
.no_proxy()
.redirect(Policy::none())
.timeout(Duration::from_secs(2))
.build()
.unwrap();
let result = runtime.web_with_client(
args,
&client,
&format!("http://{address}"),
"test-key",
&AgentCancellation::default(),
);
(result, server.join().unwrap(), runtime)
}
fn http_response(status: &str, body: &str) -> String {
format!(
"HTTP/1.1 {status}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
body.len()
)
}
#[test]
fn exa_endpoints_return_references_and_reuse_cached_content() {
let body = json!({"results":[{"url":"https://1.1.1.1/page", "title":"Docs", "text":"Rust docs: cached needle"}]}).to_string();
for operation in ["search", "open"] {
let (outcome, request, runtime) = mock_call(operation, http_response("200 OK", &body));
let result = outcome.unwrap();
assert!(result.success, "{}", result.content);
let endpoint = if operation == "search" {
"search"
} else {
"contents"
};
assert!(request.starts_with(&format!("POST /{endpoint} HTTP/1.1")));
assert!(request.contains("x-api-key: test-key"));
let body: Value =
serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap();
if operation == "search" {
assert_eq!(body["query"], "Rust docs");
assert_eq!(body["numResults"], 2);
assert_eq!(body["contents"]["text"]["maxCharacters"], 10000);
assert_eq!(body["includeDomains"], json!(["example.com"]));
assert_eq!(body["excludeDomains"], json!(["old.example.com"]));
let date = chrono::NaiveDate::parse_from_str(
body["startPublishedDate"].as_str().unwrap(),
"%Y-%m-%d",
)
.unwrap();
assert!((7..=8).contains(&(Utc::now().date_naive() - date).num_days()));
} else {
assert_eq!(body["urls"], json!(["https://1.1.1.1/page"]));
assert_eq!(body["text"]["maxCharacters"], 10000);
for key in ["includeDomains", "excludeDomains", "startPublishedDate"] {
assert!(body.get(key).is_none());
}
}
let reference = runtime.web_cache.lock().unwrap().pages[0].reference.clone();
assert!(result.content.contains(&reference));
let found = runtime.dispatch(
"web",
json!({"operation":"find", "reference":reference, "query":"needle"}),
);
assert!(found.success);
assert_eq!(found.metadata["matches"], 1);
}
}
#[test]
fn exa_errors_bounds_and_unsafe_results_fail_safely() {
for status in ["401 Unauthorized", "429 Too Many Requests", "302 Found"] {
let (outcome, _, _) =
mock_call("search", http_response(status, "secret provider body"));
let result = outcome.unwrap();
assert!(!result.success);
assert!(!result.content.contains("secret provider body"));
assert_eq!(result.metadata["rate_limited"], status.starts_with("429"));
}
let (outcome, _, _) = mock_call("open", http_response("200 OK", "not-json"));
assert!(outcome.unwrap_err().to_string().contains("malformed JSON"));
let response = format!(
"HTTP/1.1 200 OK\r\nContent-Length: {}\r\n\r\n",
RESPONSE_MAX_BYTES + 1
);
assert!(
mock_call("open", response)
.0
.unwrap_err()
.to_string()
.contains("byte limit")
);
let body = json!({"results":[{"url":"http://127.0.0.1/private", "text":"private marker"}]})
.to_string();
let (outcome, _, runtime) = mock_call("open", http_response("200 OK", &body));
let result = outcome.unwrap();
assert!(!result.success);
assert!(!result.content.contains("private marker"));
assert!(runtime.web_cache.lock().unwrap().pages.is_empty());
let (outcome, _, _) = mock_call(
"open",
http_response(
"200 OK",
r#"{"results":[],"statuses":[{"status":"error"}]}"#,
),
);
assert!(!outcome.unwrap().success);
}
#[test]
fn cache_redacts_bounds_and_cancellation_prevents_access() {
let mut cache = WebCache::default();
let text = format!(
"EXA_API_KEY=sk-testSecret123456\n{}",
"界".repeat(PAGE_MAX_BYTES)
);
let page = cache.insert("https://1.1.1.1/page".into(), "title".into(), text);
assert!(page.truncated);
assert!(page.text.len() <= PAGE_MAX_BYTES);
assert!(!page.text.contains("sk-testSecret123456"));
let output = page_result(&page);
assert!(output.content.len() <= OUTPUT_MAX_BYTES);
let temp = tempfile::TempDir::new().unwrap();
let runtime = ToolRuntime::new(temp.path()).unwrap();
let (cancellation, handle) = AgentCancellation::default().child_token();
handle.cancel();
let args =
serde_json::from_value::<WebArgs>(json!({"operation":"open","reference":"missing"}))
.unwrap();
assert_eq!(
runtime.web(args, &cancellation).unwrap_err().to_string(),
"prompt canceled"
);
}
#[test]
fn cached_find_paginates_unicode_and_misses_explicitly() {
let mut cache = WebCache::default();
let page = cache.insert(
"https://example.com".into(),
"title".into(),
"é needle one\nneedle two".into(),
);
let first = find_page(&page, "needle", 1, 0);
assert!(first.content.contains("needle one"));
assert_eq!(first.metadata["next_offset"], 1);
let second = find_page(&page, "needle", 1, 1);
assert_eq!(second.metadata["has_more"], false);
assert_eq!(find_page(&page, "absent", 5, 0).metadata["matches"], 0);
assert!(
cache
.get("unknown")
.unwrap_err()
.to_string()
.contains("cache miss")
);
assert_eq!(
cache
.insert(page.url.clone(), page.title.clone(), page.text.clone())
.reference,
page.reference
);
for n in 0..CACHE_MAX_PAGES {
cache.insert(
format!("https://example.com/{n}"),
String::new(),
"text".into(),
);
}
assert!(cache.get(&page.reference).is_err());
}
#[test]
fn fresh_duplicate_reference_survives_full_cache_insertion() {
let mut cache = WebCache::default();
let pages: Vec<_> = (0..CACHE_MAX_PAGES)
.map(|n| cache.insert(format!("https://1.1.1.1/{n}"), String::new(), "text".into()))
.collect();
let oldest = &pages[0];
let duplicate = cache.insert(
oldest.url.clone(),
oldest.title.clone(),
oldest.text.clone(),
);
let fresh = cache.insert("https://1.1.1.1/new".into(), String::new(), "new".into());
assert_eq!(duplicate.reference, oldest.reference);
assert_eq!(cache.get(&duplicate.reference).unwrap().text, oldest.text);
assert_eq!(cache.get(&fresh.reference).unwrap().text, "new");
assert!(cache.get(&pages[1].reference).is_err());
assert_eq!(cache.pages.len(), CACHE_MAX_PAGES);
}
#[test]
fn operation_validation_rejects_mixed_arguments() {
for value in [
json!({"operation":"find", "query":"x", "url":"https://example.com"}),
json!({"operation":"open","url":"https://example.com","reference":"x"}),
json!({"operation":"search","query":" ","limit":0}),
] {
assert!(
serde_json::from_value::<WebArgs>(value)
.unwrap()
.validate()
.is_err()
);
}
}
#[test]
fn public_url_rejects_credentials_and_private_networks() {
for url in [
"http://127.0.0.1",
"http://10.0.0.1",
"https://user:pass@example.com",
"file:///tmp/test",
"https://1.1.1.1/?%74oken=opaqueSecret123",
"https://1.1.1.1/?safe=yes&%61pi%5fkey=opaqueSecret123",
"https://1.1.1.1/#%74oken=opaqueSecret123",
"https://1.1.1.1/#safe=yes&%61ccess_token=opaqueSecret123",
] {
assert!(validate_url(url, &AgentCancellation::default()).is_err());
}
assert!(
validate_url(
"https://1.1.1.1/?%71=docs&input_tokens=10#section",
&AgentCancellation::default()
)
.is_ok()
);
}
#[test]
fn cached_runtime_calls_need_no_api_key_or_network() {
let temp = tempfile::TempDir::new().unwrap();
let runtime = ToolRuntime::new(temp.path()).unwrap();
let page = runtime.web_cache.lock().unwrap().insert(
"https://example.com".into(),
"title".into(),
"find this cached text".into(),
);
let cloned = runtime.clone();
let found = cloned.dispatch(
"web",
json!({"operation":"find","reference":page.reference,"query":"cached"}),
);
assert!(found.success, "{}", found.content);
assert_eq!(found.metadata["matches"], 1);
assert!(
cloned
.dispatch(
"web",
json!({"operation":"open","reference":page.reference})
)
.content
.contains("cached text")
);
}
}