diff --git a/packages/core/src/tools/web-fetch.ts b/packages/core/src/tools/web-fetch.ts index d468064f237..c6224441e38 100644 --- a/packages/core/src/tools/web-fetch.ts +++ b/packages/core/src/tools/web-fetch.ts @@ -19,7 +19,8 @@ import type { MessageBus } from '../confirmation-bus/message-bus.js'; import { ToolErrorType } from './tool-error.js'; import { getErrorMessage } from '../utils/errors.js'; import { getResponseText } from '../utils/partUtils.js'; -import { fetchWithTimeout, isPrivateIp } from '../utils/fetch.js'; +//Updated references from fetchWithTimeout to the new fetchWithSafeDns function in both the fallback (executeFallbackForUrl) and experimental (executeExperimental) fetch code paths. +import { isPrivateIpAsync, fetchWithSafeDns } from '../utils/fetch.js'; import { truncateString, wrapUntrusted } from '../utils/textUtils.js'; import { convert } from 'html-to-text'; import { @@ -267,14 +268,19 @@ class WebFetchToolInvocation extends BaseToolInvocation< ); } - private isBlockedHost(urlStr: string): boolean { + private async isBlockedHost(urlStr: string): Promise { try { const url = new URL(urlStr); - const hostname = url.hostname.toLowerCase(); - if (hostname === 'localhost' || hostname === '127.0.0.1') { + const hostname = url.hostname.toLowerCase().replace(/^\[|\]$/g, ''); + if ( + hostname === 'localhost' || + hostname === '127.0.0.1' || + hostname === '::1' || + hostname === '::ffff:127.0.0.1' + ) { return true; } - return isPrivateIp(urlStr); + return await isPrivateIpAsync(urlStr); } catch { return true; } @@ -285,7 +291,8 @@ class WebFetchToolInvocation extends BaseToolInvocation< signal: AbortSignal, ): Promise { const url = convertGithubUrlToRaw(urlStr); - if (this.isBlockedHost(url)) { + + if (await this.isBlockedHost(url)) { debugLogger.warn(`[WebFetchTool] Blocked access to host: ${url}`); throw new Error( `Access to blocked or private host ${url} is not allowed.`, @@ -294,7 +301,7 @@ class WebFetchToolInvocation extends BaseToolInvocation< const response = await retryWithBackoff( async () => { - const res = await fetchWithTimeout(url, URL_FETCH_TIMEOUT_MS, { + const res = await fetchWithSafeDns(url, URL_FETCH_TIMEOUT_MS, { signal, headers: { 'User-Agent': USER_AGENT, @@ -350,16 +357,26 @@ class WebFetchToolInvocation extends BaseToolInvocation< return textContent; } - private filterAndValidateUrls(urls: string[]): { + private async filterAndValidateUrls(urls: string[]): Promise<{ toFetch: string[]; skipped: string[]; - } { + }> { const uniqueUrls = [...new Set(urls.map(normalizeUrl))]; const toFetch: string[] = []; const skipped: string[] = []; - for (const url of uniqueUrls) { - if (this.isBlockedHost(url)) { + // Resolve blocked-host checks in parallel rather than sequentially, + // since each check can involve a DNS lookup and up to 20 URLs may + // be present in a single prompt. + const blockedChecks = await Promise.all( + uniqueUrls.map(async (url) => ({ + url, + blocked: await this.isBlockedHost(url), + })), + ); + + for (const { url, blocked } of blockedChecks) { + if (blocked) { debugLogger.warn( `[WebFetchTool] Skipped private or local host: ${url}`, ); @@ -615,7 +632,7 @@ ${aggregatedContent} // Convert GitHub blob URL to raw URL url = convertGithubUrlToRaw(url); - if (this.isBlockedHost(url)) { + if (await this.isBlockedHost(url)) { const errorMessage = `Access to blocked or private host ${url} is not allowed.`; debugLogger.warn( `[WebFetchTool] Blocked experimental fetch to host: ${url}`, @@ -633,7 +650,7 @@ ${aggregatedContent} try { const response = await retryWithBackoff( async () => { - const res = await fetchWithTimeout(url, URL_FETCH_TIMEOUT_MS, { + const res = await fetchWithSafeDns(url, URL_FETCH_TIMEOUT_MS, { signal, headers: { Accept: @@ -769,7 +786,7 @@ Response: ${rawResponseText}`; const userPrompt = this.params.prompt!; const { validUrls } = parsePrompt(userPrompt); - const { toFetch, skipped } = this.filterAndValidateUrls(validUrls); + const { toFetch, skipped } = await this.filterAndValidateUrls(validUrls); // If everything was skipped, fail early if (toFetch.length === 0 && skipped.length > 0) { diff --git a/packages/core/src/utils/fetch.ts b/packages/core/src/utils/fetch.ts index e0c34cf4f69..3a399929eeb 100644 --- a/packages/core/src/utils/fetch.ts +++ b/packages/core/src/utils/fetch.ts @@ -9,6 +9,7 @@ import { URL } from 'node:url'; import { Agent, EnvHttpProxyAgent, setGlobalDispatcher } from 'undici'; import ipaddr from 'ipaddr.js'; import { lookup } from 'node:dns/promises'; +import dns from 'node:dns'; export class FetchError extends Error { constructor( @@ -40,6 +41,8 @@ setGlobalDispatcher( }), ); +let safeDispatcher: Agent | EnvHttpProxyAgent | undefined = undefined; + export function updateGlobalFetchTimeouts(timeoutMs: number) { if (!Number.isFinite(timeoutMs) || timeoutMs <= 0) { throw new RangeError( @@ -58,6 +61,7 @@ export function updateGlobalFetchTimeouts(timeoutMs: number) { }), ); } + safeDispatcher = undefined; } /** @@ -246,4 +250,120 @@ export function setGlobalProxy(proxy: string) { bodyTimeout: defaultBodyTimeout, }), ); + /* +1. Created a secure lookup wrapper safeDnsLookup that resolves hostnames and enforces isAddressPrivate checks on the resolved addresses. + +2. Introduced a function getSafeDispatcher() that creates/caches an undici.Agent or undici.EnvHttpProxyAgent configured with the connect.lookup option pointing to safeDnsLookup. + +3. Exposed a fetchWithSafeDns() function that uses this safe dispatcher. + +4. Updated setGlobalProxy to clear/recreate the safe dispatcher when proxy configurations are updated */ + safeDispatcher = undefined; +} + +export function safeDnsLookup( + hostname: string, + options: dns.LookupOptions, + callback: ( + err: Error | null, + address: string | dns.LookupAddress[], + family?: number, + ) => void, +): void { + dns.lookup(hostname, options, (err, address, family) => { + if (err) { + return callback(err, address, family); + } + + if (Array.isArray(address)) { + const hasPrivate = address.some( + (addr) => addr && addr.address && isAddressPrivate(addr.address), + ); + if (hasPrivate) { + return callback(new PrivateIpError(), []); + } + return callback(null, address); + } + + if (typeof address === 'string' && isAddressPrivate(address)) { + return callback(new PrivateIpError(), '', family); + } + + callback(null, address, family); + }); +} +function getSafeDispatcher() { + if (safeDispatcher) { + return safeDispatcher; + } + /* [SECURITY NOTE]: When a proxy is configured via EnvHttpProxyAgent, + safeDnsLookup protects direct connections. Proxied CONNECT requests + resolve DNS on the proxy server itself. */ + const connectOptions = { + lookup: safeDnsLookup, + }; + + if (currentProxy) { + const noProxy = ( + process.env['NO_PROXY'] ?? + process.env['no_proxy'] ?? + '' + )?.trim(); + safeDispatcher = new EnvHttpProxyAgent({ + httpProxy: currentProxy, + httpsProxy: currentProxy, + noProxy, + headersTimeout: defaultHeadersTimeout, + bodyTimeout: defaultBodyTimeout, + connect: connectOptions, + }); + } else { + safeDispatcher = new Agent({ + headersTimeout: defaultHeadersTimeout, + bodyTimeout: defaultBodyTimeout, + connect: connectOptions, + }); + } + + return safeDispatcher; +} + +export async function fetchWithSafeDns( + url: string, + timeout: number, + options?: RequestInit, +): Promise { + const controller = new AbortController(); + const timeoutId = setTimeout(() => controller.abort(), timeout); + + if (options?.signal) { + if (options.signal.aborted) { + controller.abort(); + } else { + options.signal.addEventListener('abort', () => controller.abort(), { + once: true, + }); + } + } + + try { + const dispatcher = getSafeDispatcher(); + const response = await fetch(url, { + ...options, + signal: controller.signal, + // @ts-expect-error dispatcher is supported by undici-backed fetch in Node.js + dispatcher, + }); + return response; + } catch (error) { + if (isAbortError(error)) { + if (options?.signal?.aborted) { + throw options.signal.reason ?? error; + } + throw new FetchError(`Request timed out after ${timeout}ms`, 'ETIMEDOUT'); + } + throw new FetchError(getErrorMessage(error), undefined, { cause: error }); + } finally { + clearTimeout(timeoutId); + } }