fix: serialize pinned dns web fetches

This commit is contained in:
hamb1y
2026-07-07 15:40:53 +08:00
committed by Xubin Ren
parent 4353f4680b
commit 97e3b360c2
3 changed files with 76 additions and 31 deletions
+15 -12
View File
@@ -118,6 +118,12 @@ def _resolve_url_safe(url: str) -> tuple[bool, str, tuple[str, ...]]:
return resolve_url_target(url)
def _pinned_dns_transport(proxy: str | None = None) -> httpx.AsyncBaseTransport:
from nanobot.security.network import PinnedDNSAsyncTransport
return PinnedDNSAsyncTransport(proxy=proxy)
async def _get_with_safe_redirects(
client: httpx.AsyncClient,
url: str,
@@ -126,14 +132,11 @@ async def _get_with_safe_redirects(
"""GET a URL while validating every redirect target before requesting it."""
current_url = url
for _ in range(MAX_REDIRECTS + 1):
is_valid, error_msg, resolved_ips = _resolve_url_safe(current_url)
is_valid, error_msg, _ = _resolve_url_safe(current_url)
if not is_valid:
return None, f"Redirect blocked: {error_msg}"
from nanobot.security.network import pin_resolved_url_dns
with pin_resolved_url_dns(current_url, resolved_ips):
response = await client.get(current_url, headers=headers, follow_redirects=False)
response = await client.get(current_url, headers=headers, follow_redirects=False)
is_redirect = 300 <= response.status_code < 400
if not is_redirect:
return response, None
@@ -162,20 +165,17 @@ async def _stream_with_safe_redirects(
"""Open a streamed response while validating every redirect target first."""
current_url = url
for _ in range(MAX_REDIRECTS + 1):
is_valid, error_msg, resolved_ips = _resolve_url_safe(current_url)
is_valid, error_msg, _ = _resolve_url_safe(current_url)
if not is_valid:
return None, None, f"Redirect blocked: {error_msg}"
from nanobot.security.network import pin_resolved_url_dns
stream = client.stream(
"GET",
current_url,
headers=headers,
follow_redirects=False,
)
with pin_resolved_url_dns(current_url, resolved_ips):
response = await stream.__aenter__()
response = await stream.__aenter__()
is_redirect = 300 <= response.status_code < 400
if not is_redirect:
return response, stream, None
@@ -966,7 +966,10 @@ class WebFetchTool(Tool):
# Detect and fetch images directly to avoid Jina's textual image captioning
try:
async with httpx.AsyncClient(proxy=self.proxy, timeout=15.0) as client:
async with httpx.AsyncClient(
transport=_pinned_dns_transport(self.proxy),
timeout=15.0,
) as client:
r, stream, redirect_error = await _stream_with_safe_redirects(
client,
url,
@@ -1037,7 +1040,7 @@ class WebFetchTool(Tool):
try:
async with httpx.AsyncClient(
timeout=30.0,
proxy=self.proxy,
transport=_pinned_dns_transport(self.proxy),
) as client:
r, redirect_error = await _get_with_safe_redirects(
client,
+20 -5
View File
@@ -2,6 +2,7 @@
from __future__ import annotations
import asyncio
import ipaddress
import re
import socket
@@ -114,7 +115,12 @@ def validate_url_target(url: str, *, allow_loopback: bool = False) -> tuple[bool
@contextmanager
def pin_resolved_url_dns(url: str, resolved_ips: tuple[str, ...]):
"""Pin DNS lookups for the URL hostname to previously validated IPs."""
"""Pin DNS lookups for the URL hostname to previously validated IPs.
This temporarily overrides process-global resolver state. Do not use it
directly across awaits unless the caller serializes access; prefer
PinnedDNSAsyncTransport for HTTP requests.
"""
try:
hostname = urlparse(url).hostname
except Exception:
@@ -149,17 +155,26 @@ def pin_resolved_url_dns(url: str, resolved_ips: tuple[str, ...]):
class PinnedDNSAsyncTransport(httpx.AsyncBaseTransport):
"""HTTPX transport that pins each request to the IPs validated for its URL."""
def __init__(self, *, allow_loopback: bool = False) -> None:
_resolver_lock = asyncio.Lock()
def __init__(
self,
*,
allow_loopback: bool = False,
proxy: httpx.ProxyTypes | None = None,
inner: httpx.AsyncBaseTransport | None = None,
) -> None:
self._allow_loopback = allow_loopback
self._inner = httpx.AsyncHTTPTransport()
self._inner = inner or httpx.AsyncHTTPTransport(proxy=proxy)
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
url = str(request.url)
ok, error, resolved_ips = resolve_url_target(url, allow_loopback=self._allow_loopback)
if not ok:
raise httpx.RequestError(error, request=request)
with pin_resolved_url_dns(url, resolved_ips):
return await self._inner.handle_async_request(request)
async with self._resolver_lock:
with pin_resolved_url_dns(url, resolved_ips):
return await self._inner.handle_async_request(request)
async def aclose(self) -> None:
await self._inner.aclose()