use super::contract::metadata_key as meta;
use crate::{cancellation::AgentCancellation, output::redact_sensitive_text};
use dom_smoothie::{Config, Readability, TextMode};
use encoding_rs::{Encoding, UTF_8};
use reqwest::{StatusCode, Url, blocking::Client, redirect::Policy};
use serde_json::json;
use std::{
io::Read,
net::{IpAddr, SocketAddr, ToSocketAddrs},
sync::{
atomic::{AtomicUsize, Ordering},
mpsc,
},
thread::{self, JoinHandle},
time::{Duration, Instant},
};
const URL_FETCH_CONNECT_TIMEOUT: Duration = Duration::from_secs(5);
const URL_FETCH_DNS_TIMEOUT: Duration = Duration::from_secs(5);
const URL_FETCH_MAX_DNS_RESOLVER_WORKERS: usize = 8;
const URL_FETCH_REQUEST_TIMEOUT: Duration = Duration::from_secs(15);
const URL_FETCH_MAX_BYTES: usize = 512 * 1024;
const URL_FETCH_OUTPUT_MAX_BYTES: usize = 48 * 1024;
const URL_FETCH_DEFAULT_MAX_TOKENS: u64 = 5_000;
const URL_FETCH_MIN_MAX_TOKENS: u64 = 1_000;
const URL_FETCH_MAX_MAX_TOKENS: u64 = 10_000;
const URL_FETCH_READABILITY_MIN_LENGTH: usize = 200;
const TRUNCATION_MARKER: &str = "\n[truncated]";
const URL_FETCH_CANCEL_POLL_INTERVAL: Duration = Duration::from_millis(25);
const URL_FETCH_WORKER_JOIN_TIMEOUT: Duration = Duration::from_millis(500);
static URL_FETCH_DNS_RESOLVER_WORKER_COUNT: AtomicUsize = AtomicUsize::new(0);
#[derive(Debug, Clone)]
struct PinnedTarget {
host: String,
addrs: Vec<SocketAddr>,
}
#[derive(Debug, Clone)]
struct ExtractedContent {
text: String,
title: Option<String>,
}
#[derive(Debug, Clone)]
pub(crate) struct UrlFetchInput {
pub(crate) url: String,
pub(crate) max_tokens: Option<u64>,
}
impl UrlFetchInput {
pub(crate) fn validate(mut self) -> anyhow::Result<Self> {
self.url = self.url.trim().to_string();
if self.url.is_empty() {
anyhow::bail!("url must not be empty");
}
let parsed = Url::parse(&self.url)
.map_err(|_| anyhow::anyhow!("url must be an http(s) URL with a host"))?;
if !matches!(parsed.scheme(), "http" | "https") || parsed.host_str().is_none() {
anyhow::bail!("url must be an http(s) URL with a host");
}
if let Some(max_tokens) = self.max_tokens
&& !(URL_FETCH_MIN_MAX_TOKENS..=URL_FETCH_MAX_MAX_TOKENS).contains(&max_tokens)
{
anyhow::bail!(
"maxTokens must be between {URL_FETCH_MIN_MAX_TOKENS} and {URL_FETCH_MAX_MAX_TOKENS}"
);
}
Ok(self)
}
fn max_tokens(&self) -> u64 {
self.max_tokens.unwrap_or(URL_FETCH_DEFAULT_MAX_TOKENS)
}
}
#[derive(Debug, Clone)]
pub(crate) struct UrlFetchOutput {
pub(crate) success: bool,
pub(crate) content: String,
pub(crate) metadata: serde_json::Value,
}
pub(crate) fn fetch_url(
input: UrlFetchInput,
cancellation: &AgentCancellation,
) -> anyhow::Result<UrlFetchOutput> {
cancellation.check()?;
let url = Url::parse(&input.url)?;
let response = fetch_protected_response(url.clone(), cancellation)?;
fetch_with_response(input, url, response, cancellation)
}
fn build_pinned_client(target: &PinnedTarget) -> anyhow::Result<Client> {
Client::builder()
.connect_timeout(URL_FETCH_CONNECT_TIMEOUT)
.timeout(URL_FETCH_REQUEST_TIMEOUT)
.user_agent(format!("magi-code/{}", env!("CARGO_PKG_VERSION")))
.redirect(Policy::none())
.no_proxy()
.resolve_to_addrs(&target.host, &target.addrs)
.build()
.map_err(|error| anyhow::anyhow!("url_fetch HTTP client setup failed: {error}"))
}
fn fetch_protected_response(
url: Url,
cancellation: &AgentCancellation,
) -> anyhow::Result<FetchHttpResponse> {
cancellation.check()?;
let target = resolve_public_target(&url, cancellation)?;
cancellation.check()?;
let client = build_pinned_client(&target)?;
cancellation.check()?;
fetch_response_with_client(&client, &url, cancellation)
}
#[derive(Debug)]
struct FetchHttpResponse {
status: StatusCode,
content_type: String,
bytes: Vec<u8>,
}
fn is_html_content_type(content_type: &str) -> bool {
let media_type = content_type.split(';').next().map(str::trim);
media_type.is_some_and(|media_type| {
media_type.eq_ignore_ascii_case("text/html")
|| media_type.eq_ignore_ascii_case("application/xhtml+xml")
})
}
fn decode_response_body(bytes: &[u8], content_type: &str) -> String {
if let Some((encoding, _)) = Encoding::for_bom(bytes) {
return encoding.decode(bytes).0.into_owned();
}
if let Some(encoding) =
declared_charset(content_type).and_then(|label| Encoding::for_label(label.as_bytes()))
{
return encoding.decode_without_bom_handling(bytes).0.into_owned();
}
if is_html_content_type(content_type)
&& let Some(label) = sniff_html_meta_charset(bytes)
&& let Some(encoding) = meta_charset_encoding(&label)
{
return encoding.decode_without_bom_handling(bytes).0.into_owned();
}
UTF_8.decode_without_bom_handling(bytes).0.into_owned()
}
fn declared_charset(content_type: &str) -> Option<String> {
content_type.split(';').skip(1).find_map(|parameter| {
let (name, value) = parameter.split_once('=')?;
if !name.trim().eq_ignore_ascii_case("charset") {
return None;
}
let value = value.trim().trim_matches(['"', '\'']).trim();
(!value.is_empty()).then(|| value.to_string())
})
}
fn sniff_html_meta_charset(bytes: &[u8]) -> Option<String> {
let prefix = &bytes[..bytes.len().min(1024)];
let lower = prefix
.iter()
.map(|byte| byte.to_ascii_lowercase())
.collect::<Vec<_>>();
let mut cursor = 0;
while cursor + b"<meta".len() <= lower.len() {
if lower[cursor..].starts_with(b"<meta")
&& lower
.get(cursor + b"<meta".len())
.is_none_or(|byte| byte.is_ascii_whitespace() || *byte == b'/' || *byte == b'>')
{
let attributes_start = cursor + b"<meta".len();
let tag_end = lower[attributes_start..]
.iter()
.position(|byte| *byte == b'>')
.map_or(lower.len(), |offset| attributes_start + offset);
if let Some(label) = charset_value(&lower[attributes_start..tag_end]) {
return Some(label);
}
cursor = tag_end.saturating_add(1);
} else {
cursor += 1;
}
}
None
}
fn charset_value(attributes: &[u8]) -> Option<String> {
let mut cursor = 0;
while cursor + b"charset".len() <= attributes.len() {
if attributes[cursor..].starts_with(b"charset")
&& (cursor == 0 || !is_attribute_name_byte(attributes[cursor - 1]))
&& (cursor + b"charset".len() == attributes.len()
|| !is_attribute_name_byte(attributes[cursor + b"charset".len()]))
{
let mut value_start = cursor + b"charset".len();
while attributes
.get(value_start)
.is_some_and(u8::is_ascii_whitespace)
{
value_start += 1;
}
if attributes.get(value_start) == Some(&b'=') {
value_start += 1;
while attributes
.get(value_start)
.is_some_and(u8::is_ascii_whitespace)
{
value_start += 1;
}
let quote = attributes.get(value_start).copied();
let (value_start, value_end) = match quote {
Some(quote @ (b'"' | b'\'')) => {
let value_start = value_start + 1;
let value_end = attributes[value_start..]
.iter()
.position(|byte| *byte == quote)
.map_or(attributes.len(), |offset| value_start + offset);
(value_start, value_end)
}
_ => {
let value_end = attributes[value_start..]
.iter()
.position(|byte| {
byte.is_ascii_whitespace() || matches!(*byte, b';' | b'/')
})
.map_or(attributes.len(), |offset| value_start + offset);
(value_start, value_end)
}
};
if value_start < value_end {
return Some(
String::from_utf8_lossy(&attributes[value_start..value_end]).into_owned(),
);
}
}
}
cursor += 1;
}
None
}
fn is_attribute_name_byte(byte: u8) -> bool {
byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b':')
}
fn meta_charset_encoding(label: &str) -> Option<&'static Encoding> {
if label.eq_ignore_ascii_case("utf-16")
|| label.eq_ignore_ascii_case("utf-16le")
|| label.eq_ignore_ascii_case("utf-16be")
{
return Some(UTF_8);
}
Encoding::for_label(label.as_bytes())
}
struct DnsResolverPermit<'a> {
counter: &'a AtomicUsize,
}
impl<'a> DnsResolverPermit<'a> {
fn try_acquire(counter: &'a AtomicUsize, limit: usize) -> Option<Self> {
if limit == 0 {
return None;
}
let mut current = counter.load(Ordering::Acquire);
loop {
if current >= limit {
return None;
}
match counter.compare_exchange_weak(
current,
current + 1,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => return Some(Self { counter }),
Err(observed) => current = observed,
}
}
}
}
impl Drop for DnsResolverPermit<'_> {
fn drop(&mut self) {
self.counter.fetch_sub(1, Ordering::Release);
}
}
fn acquire_dns_resolver_permit<'a>(
counter: &'a AtomicUsize,
limit: usize,
) -> anyhow::Result<DnsResolverPermit<'a>> {
DnsResolverPermit::try_acquire(counter, limit)
.ok_or_else(|| anyhow::anyhow!("url_fetch DNS resolver capacity exhausted; retry later"))
}
struct UrlFetchWorkerHandle {
worker_label: &'static str,
done_receiver: mpsc::Receiver<()>,
join_handle: JoinHandle<()>,
}
impl UrlFetchWorkerHandle {
fn join_or_warn(self) {
match self
.done_receiver
.recv_timeout(URL_FETCH_WORKER_JOIN_TIMEOUT)
{
Ok(()) | Err(mpsc::RecvTimeoutError::Disconnected) => {
if self.join_handle.join().is_err() {
eprintln!(
"magi-code warning: url_fetch {} worker panicked",
self.worker_label
);
}
}
Err(mpsc::RecvTimeoutError::Timeout) => {
eprintln!(
"magi-code warning: url_fetch {} worker did not exit within cleanup grace period; detaching worker",
self.worker_label
);
}
}
}
fn detach_with_warning(self) {
let detail = match self.worker_label {
"DNS" => {
"worker detached; the OS resolver may continue running after the caller deadline"
}
_ => "worker detached; it may continue until its own timeout",
};
eprintln!(
"magi-code warning: url_fetch {} {}",
self.worker_label, detail
);
let UrlFetchWorkerHandle { join_handle, .. } = self;
drop(join_handle);
}
}
fn spawn_url_fetch_worker(
client: Client,
url: Url,
) -> anyhow::Result<(
mpsc::Receiver<anyhow::Result<FetchHttpResponse>>,
UrlFetchWorkerHandle,
)> {
let (sender, receiver) = mpsc::sync_channel(1);
let (done_sender, done_receiver) = mpsc::sync_channel(1);
let join_handle = thread::Builder::new()
.name("url-fetch-http".to_string())
.spawn(move || {
let _ = sender.send(blocking_fetch_response(client, url));
let _ = done_sender.send(());
})?;
Ok((
receiver,
UrlFetchWorkerHandle {
worker_label: "HTTP",
done_receiver,
join_handle,
},
))
}
fn spawn_dns_resolution_worker(
url: Url,
) -> anyhow::Result<(
mpsc::Receiver<anyhow::Result<PinnedTarget>>,
UrlFetchWorkerHandle,
)> {
let permit = acquire_dns_resolver_permit(
&URL_FETCH_DNS_RESOLVER_WORKER_COUNT,
URL_FETCH_MAX_DNS_RESOLVER_WORKERS,
)?;
let (sender, receiver) = mpsc::sync_channel(1);
let (done_sender, done_receiver) = mpsc::sync_channel(1);
let join_handle = thread::Builder::new()
.name("url-fetch-dns".to_string())
.spawn(move || {
let _permit = permit;
let _ = sender.send(resolve_public_target_sync(&url));
let _ = done_sender.send(());
})?;
Ok((
receiver,
UrlFetchWorkerHandle {
worker_label: "DNS",
done_receiver,
join_handle,
},
))
}
fn recv_cancellable<T>(
receiver: &mpsc::Receiver<T>,
timeout_label: &str,
idle_timeout: Duration,
cancellation: &AgentCancellation,
) -> anyhow::Result<T> {
let start = Instant::now();
loop {
cancellation.check()?;
if start.elapsed() >= idle_timeout {
anyhow::bail!("{timeout_label} after {}s", idle_timeout.as_secs());
}
let remaining = idle_timeout.saturating_sub(start.elapsed());
match receiver.recv_timeout(remaining.min(URL_FETCH_CANCEL_POLL_INTERVAL)) {
Ok(value) => return Ok(value),
Err(mpsc::RecvTimeoutError::Timeout) => continue,
Err(mpsc::RecvTimeoutError::Disconnected) => {
cancellation.check()?;
anyhow::bail!("url_fetch worker disconnected");
}
}
}
}
fn wait_for_worker_result<T>(
receiver: mpsc::Receiver<anyhow::Result<T>>,
worker: UrlFetchWorkerHandle,
timeout_label: &str,
timeout: Duration,
cancellation: &AgentCancellation,
) -> anyhow::Result<T> {
let result = match recv_cancellable(&receiver, timeout_label, timeout, cancellation) {
Ok(result) => result,
Err(error) => {
drop(receiver);
worker.detach_with_warning();
return Err(error);
}
};
worker.join_or_warn();
result
}
fn blocking_fetch_response(client: Client, url: Url) -> anyhow::Result<FetchHttpResponse> {
let response = client.get(url).send().map_err(|error| {
if error.is_timeout() {
anyhow::anyhow!("url_fetch request timed out")
} else {
anyhow::anyhow!("url_fetch network error: {error}")
}
})?;
let status = response.status();
let content_type = response
.headers()
.get(reqwest::header::CONTENT_TYPE)
.and_then(|value| value.to_str().ok())
.unwrap_or_default()
.to_string();
if response
.content_length()
.is_some_and(|length| length > URL_FETCH_MAX_BYTES as u64)
{
anyhow::bail!("url_fetch response exceeded {URL_FETCH_MAX_BYTES} byte limit");
}
let bytes = read_bounded_body(response)?;
Ok(FetchHttpResponse {
status,
content_type,
bytes,
})
}
fn fetch_response_with_client(
client: &Client,
url: &Url,
cancellation: &AgentCancellation,
) -> anyhow::Result<FetchHttpResponse> {
let (receiver, worker) = spawn_url_fetch_worker(client.clone(), url.clone())?;
let response = wait_for_worker_result(
receiver,
worker,
"url_fetch response timeout",
URL_FETCH_REQUEST_TIMEOUT,
cancellation,
)?;
cancellation.check()?;
Ok(response)
}
fn fetch_with_response(
args: UrlFetchInput,
url: Url,
response: FetchHttpResponse,
cancellation: &AgentCancellation,
) -> anyhow::Result<UrlFetchOutput> {
let status = response.status;
let content_type = response.content_type;
let byte_count = response.bytes.len();
let raw = decode_response_body(&response.bytes, &content_type);
if !status.is_success() {
let (content, truncated) = sanitize_and_bound(&raw, args.max_tokens());
return Ok(result(FetchResultParts {
success: false,
content,
url: &url,
status,
content_type: &content_type,
title: None,
bytes: byte_count,
truncated,
}));
}
let extracted = extract_content(&raw, &content_type, url.as_str());
cancellation.check()?;
let (content, truncated) = sanitize_and_bound(&extracted.text, args.max_tokens());
cancellation.check()?;
Ok(result(FetchResultParts {
success: true,
content,
url: &url,
status,
content_type: &content_type,
title: extracted.title.as_deref(),
bytes: byte_count,
truncated,
}))
}
struct FetchResultParts<'a> {
success: bool,
content: String,
url: &'a Url,
status: StatusCode,
content_type: &'a str,
title: Option<&'a str>,
bytes: usize,
truncated: bool,
}
fn result(parts: FetchResultParts<'_>) -> UrlFetchOutput {
UrlFetchOutput {
success: parts.success,
content: parts.content,
metadata: json!({
(meta::URL): parts.url.as_str(),
(meta::STATUS): parts.status.as_u16(),
(meta::CONTENT_TYPE): parts.content_type,
(meta::TITLE): parts.title,
(meta::BYTES): parts.bytes,
(meta::TRUNCATED): parts.truncated,
}),
}
}
fn read_bounded_body(mut response: reqwest::blocking::Response) -> anyhow::Result<Vec<u8>> {
let mut bytes = Vec::new();
response
.by_ref()
.take((URL_FETCH_MAX_BYTES + 1) as u64)
.read_to_end(&mut bytes)
.map_err(|error| anyhow::anyhow!("url_fetch response read failed: {error}"))?;
if bytes.len() > URL_FETCH_MAX_BYTES {
anyhow::bail!("url_fetch response exceeded {URL_FETCH_MAX_BYTES} byte limit");
}
Ok(bytes)
}
fn extract_content(body: &str, content_type: &str, url: &str) -> ExtractedContent {
if !content_type.to_ascii_lowercase().contains("html") {
return ExtractedContent {
text: body.to_string(),
title: None,
};
}
if let Some(extracted) = readability_extract(body, url) {
return extracted;
}
ExtractedContent {
text: html2md::rewrite_html(body, false),
title: None,
}
}
fn readability_extract(html: &str, url: &str) -> Option<ExtractedContent> {
let cfg = Config {
text_mode: TextMode::Markdown,
..Default::default()
};
let mut readability = Readability::new(html, Some(url), Some(cfg)).ok()?;
let article = readability.parse().ok()?;
let text = article.text_content.to_string();
let text = text.trim().to_string();
if article.length <= URL_FETCH_READABILITY_MIN_LENGTH || text.is_empty() {
return None;
}
let title = (!article.title.trim().is_empty()).then(|| article.title.trim().to_string());
Some(ExtractedContent { text, title })
}
fn sanitize_and_bound(text: &str, max_tokens: u64) -> (String, bool) {
let cap = URL_FETCH_OUTPUT_MAX_BYTES.min(max_tokens.saturating_mul(4) as usize);
let mut truncated = false;
let mut content = truncate_with_marker(text.trim(), cap, &mut truncated);
content = redact_sensitive_text(&content);
if content.len() > cap {
truncated = true;
content = truncate_with_marker(&content, cap, &mut truncated);
}
(content, truncated)
}
fn truncate_with_marker(text: &str, max_bytes: usize, truncated: &mut bool) -> String {
if text.len() <= max_bytes {
return text.to_string();
}
*truncated = true;
if max_bytes <= TRUNCATION_MARKER.len() {
return truncate_to_char_boundary(TRUNCATION_MARKER, max_bytes).to_string();
}
let text_budget = max_bytes - TRUNCATION_MARKER.len();
format!(
"{}{}",
truncate_to_char_boundary(text, text_budget),
TRUNCATION_MARKER
)
}
fn truncate_to_char_boundary(text: &str, max_bytes: usize) -> &str {
if text.len() <= max_bytes {
return text;
}
let mut end = max_bytes;
while end > 0 && !text.is_char_boundary(end) {
end -= 1;
}
&text[..end]
}
pub(super) fn validate_public_url(
url: &Url,
cancellation: &AgentCancellation,
) -> anyhow::Result<()> {
cancellation.check()?;
resolve_public_target(url, cancellation).map(|_| ())
}
fn resolve_public_target(
url: &Url,
cancellation: &AgentCancellation,
) -> anyhow::Result<PinnedTarget> {
let (receiver, worker) = spawn_dns_resolution_worker(url.clone())?;
wait_for_worker_result(
receiver,
worker,
"url_fetch DNS resolution timeout",
URL_FETCH_DNS_TIMEOUT,
cancellation,
)
}
fn resolve_public_target_sync(url: &Url) -> anyhow::Result<PinnedTarget> {
if !matches!(url.scheme(), "http" | "https") {
anyhow::bail!("url_fetch only supports http and https URLs");
}
let host = url
.host_str()
.ok_or_else(|| anyhow::anyhow!("url_fetch URL must include a host"))?
.to_string();
let port = url.port_or_known_default().unwrap_or(80);
let addrs = resolve_host_addrs(&host, port)?;
validate_public_addrs(&host, addrs)
}
fn resolve_host_addrs(host: &str, port: u16) -> anyhow::Result<Vec<SocketAddr>> {
if let Ok(ip) = host.parse::<IpAddr>() {
return Ok(vec![SocketAddr::new(ip, port)]);
}
(host, port)
.to_socket_addrs()
.map(|iter| iter.collect::<Vec<_>>())
.map_err(|error| anyhow::anyhow!("url_fetch DNS resolution failed for {host}: {error}"))
}
fn validate_public_addrs(host: &str, addrs: Vec<SocketAddr>) -> anyhow::Result<PinnedTarget> {
if addrs.is_empty() {
anyhow::bail!("url_fetch DNS resolution returned no addresses for {host}");
}
for addr in &addrs {
if !is_public_ip(addr.ip()) {
anyhow::bail!(
"url_fetch target resolves to unsafe IP address: {}",
addr.ip()
);
}
}
Ok(PinnedTarget {
host: host.to_string(),
addrs,
})
}
fn is_public_ipv4(ip: std::net::Ipv4Addr) -> bool {
let value = u32::from(ip);
!(value & 0xff00_0000 == 0x0000_0000 || value & 0xff00_0000 == 0x0a00_0000 || value & 0xfff0_0000 == 0xac10_0000 || value & 0xffff_0000 == 0xc0a8_0000 || value & 0xff00_0000 == 0x7f00_0000 || value & 0xffff_0000 == 0xa9fe_0000 || value & 0xffc0_0000 == 0x6440_0000 || value & 0xffff_ff00 == 0xc000_0000 || value & 0xffff_ff00 == 0xc000_0200 || value & 0xffff_ff00 == 0xc633_6400 || value & 0xffff_ff00 == 0xcb00_7100 || value & 0xffff_ff00 == 0xc058_6300 || value & 0xfffe_0000 == 0xc612_0000 || value & 0xf000_0000 == 0xe000_0000 || value & 0xf000_0000 == 0xf000_0000) }
fn is_public_ip(ip: IpAddr) -> bool {
match ip {
IpAddr::V4(ip) => is_public_ipv4(ip),
IpAddr::V6(ip) => {
if let Some(mapped) = ip.to_ipv4_mapped() {
return is_public_ipv4(mapped);
}
let segments = ip.segments();
let compatible = segments[..6].iter().all(|segment| *segment == 0);
let first = segments[0];
!(
ip.is_loopback()
|| ip.is_unspecified()
|| ip.is_multicast()
|| compatible || (first & 0xfe00) == 0xfc00 || (first & 0xffc0) == 0xfe80 || (first & 0xffc0) == 0xfec0 || first == 0x2002 || (first == 0x2001 && segments[1] <= 0x01ff) || (first == 0x2001 && segments[1] == 0x0db8) || (first == 0x3fff && segments[1] <= 0x0fff) || first == 0x5f00 || (first == 0x0064
&& segments[1] == 0xff9b
&& (segments[2..6].iter().all(|segment| *segment == 0)
|| segments[2] == 1)) || (first == 0x0100
&& segments[1..4].iter().all(|segment| *segment == 0))
)
}
}
}