Yash030/claude-code-proxy
2
1"""Outbound HTTP for web_search / web_fetch (client, body caps, logging)."""2 3from __future__ import annotations4 5import asyncio6import socket7from collections.abc import AsyncIterator8from urllib.parse import urljoin, urlparse9 10import aiohttp11import httpx12from aiohttp import ClientSession, ClientTimeout, TCPConnector13from aiohttp.abc import AbstractResolver, ResolveResult14from loguru import logger15 16from . import constants17from .constants import (18 _MAX_FETCH_CHARS,19 _MAX_SEARCH_RESULTS,20 _REDIRECT_RESPONSE_BODY_CAP_BYTES,21 _REQUEST_TIMEOUT_S,22 _WEB_FETCH_REDIRECT_STATUSES,23 _WEB_TOOL_HTTP_HEADERS,24)25from .egress import (26 WebFetchEgressPolicy,27 WebFetchEgressViolation,28 get_validated_stream_addrinfos_for_egress,29)30from .parsers import HTMLTextParser, SearchResultParser31 32 33def _safe_public_host_for_logs(url: str) -> str:34 host = urlparse(url).hostname or ""35 return host[:253]36 37 38def _log_web_tool_failure(39 tool_name: str,40 error: BaseException,41 *,42 fetch_url: str | None = None,43) -> None:44 exc_type = type(error).__name__45 if isinstance(error, WebFetchEgressViolation):46 host = _safe_public_host_for_logs(fetch_url) if fetch_url else ""47 logger.warning(48 "web_tool_egress_rejected tool={} exc_type={} host={!r}",49 tool_name,50 exc_type,51 host,52 )53 return54 if tool_name == "web_fetch" and fetch_url:55 logger.warning(56 "web_tool_failure tool={} exc_type={} host={!r}",57 tool_name,58 exc_type,59 _safe_public_host_for_logs(fetch_url),60 )61 else:62 logger.warning("web_tool_failure tool={} exc_type={}", tool_name, exc_type)63 64 65def _web_tool_client_error_summary(66 tool_name: str,67 error: BaseException,68 *,69 verbose: bool,70) -> str:71 if verbose:72 return f"{tool_name} failed: {type(error).__name__}"73 return "Web tool request failed."74 75 76async def _iter_response_body_under_cap(77 response: httpx.Response, max_bytes: int78) -> AsyncIterator[bytes]:79 if max_bytes <= 0:80 return81 received = 082 async for chunk in response.aiter_bytes(chunk_size=65_536):83 if received >= max_bytes:84 break85 remaining = max_bytes - received86 if len(chunk) <= remaining:87 received += len(chunk)88 yield chunk89 if received >= max_bytes:90 break91 else:92 yield chunk[:remaining]93 break94 95 96async def _drain_response_body_capped(response: httpx.Response, max_bytes: int) -> None:97 async for _ in _iter_response_body_under_cap(response, max_bytes):98 pass99 100 101async def _read_response_body_capped(response: httpx.Response, max_bytes: int) -> bytes:102 return b"".join(103 [piece async for piece in _iter_response_body_under_cap(response, max_bytes)]104 )105 106 107_NUMERIC_RESOLVE_FLAGS = socket.AI_NUMERICHOST | socket.AI_NUMERICSERV108_NAME_RESOLVE_FLAGS = socket.NI_NUMERICHOST | socket.NI_NUMERICSERV109 110 111def getaddrinfo_rows_to_resolve_results(112 host: str, addrinfos: list[tuple]113) -> list[ResolveResult]:114 """Map :func:`socket.getaddrinfo` rows to aiohttp :class:`ResolveResult` (ThreadedResolver logic)."""115 out: list[ResolveResult] = []116 for family, _type, proto, _canon, sockaddr in addrinfos:117 if family == socket.AF_INET6:118 if len(sockaddr) < 3:119 continue120 if sockaddr[3]:121 resolved_host, port = socket.getnameinfo(sockaddr, _NAME_RESOLVE_FLAGS)122 else:123 resolved_host, port = sockaddr[:2]124 else:125 assert family == socket.AF_INET, family126 resolved_host, port = sockaddr[0], sockaddr[1]127 resolved_host = str(resolved_host)128 port = int(port)129 out.append(130 ResolveResult(131 hostname=host,132 host=resolved_host,133 port=int(port),134 family=family,135 proto=proto,136 flags=_NUMERIC_RESOLVE_FLAGS,137 )138 )139 return out140 141 142class _PinnedEgressStaticResolver(AbstractResolver):143 """Return only pre-validated :class:`ResolveResult` for the outbound request."""144 145 def __init__(self, results: list[ResolveResult]) -> None:146 self._results = results147 148 async def resolve(149 self, host: str, port: int = 0, family: int = socket.AF_INET150 ) -> list[ResolveResult]:151 return self._results152 153 async def close(self) -> None: # pragma: no cover - aiohttp contract154 return155 156 157async def _read_aiohttp_body_capped(158 response: aiohttp.ClientResponse, max_bytes: int159) -> bytes:160 received = 0161 parts: list[bytes] = []162 async for chunk in response.content.iter_chunked(65_536):163 if received >= max_bytes:164 break165 remaining = max_bytes - received166 if len(chunk) <= remaining:167 received += len(chunk)168 parts.append(chunk)169 else:170 parts.append(chunk[:remaining])171 break172 return b"".join(parts)173 174 175async def _drain_aiohttp_body_capped(176 response: aiohttp.ClientResponse, max_bytes: int177) -> None:178 if max_bytes <= 0:179 return180 received = 0181 async for chunk in response.content.iter_chunked(65_536):182 received += len(chunk)183 if received >= max_bytes:184 break185 186 187async def _run_web_search(query: str) -> list[dict[str, str]]:188 async with (189 httpx.AsyncClient(190 timeout=_REQUEST_TIMEOUT_S,191 follow_redirects=True,192 headers=_WEB_TOOL_HTTP_HEADERS,193 ) as client,194 client.stream(195 "GET",196 "https://lite.duckduckgo.com/lite/",197 params={"q": query},198 ) as response,199 ):200 response.raise_for_status()201 body_bytes = await _read_response_body_capped(202 response, constants._MAX_WEB_FETCH_RESPONSE_BYTES203 )204 text = body_bytes.decode("utf-8", errors="replace")205 parser = SearchResultParser()206 parser.feed(text)207 return parser.results[:_MAX_SEARCH_RESULTS]208 209 210async def _run_web_fetch(url: str, egress: WebFetchEgressPolicy) -> dict[str, str]:211 """Fetch URL with manual redirects; each hop is DNS-pinned to validated addresses."""212 current_url = url213 redirect_hops = 0214 timeout = ClientTimeout(total=_REQUEST_TIMEOUT_S)215 216 while True:217 addr_infos = await asyncio.to_thread(218 get_validated_stream_addrinfos_for_egress, current_url, egress219 )220 host = urlparse(current_url).hostname or ""221 results = getaddrinfo_rows_to_resolve_results(host, addr_infos)222 resolver = _PinnedEgressStaticResolver(results)223 connector = TCPConnector(224 resolver=resolver,225 force_close=True,226 )227 try:228 async with (229 ClientSession(230 timeout=timeout,231 headers=_WEB_TOOL_HTTP_HEADERS,232 connector=connector,233 ) as session,234 session.get(current_url, allow_redirects=False) as response,235 ):236 if response.status in _WEB_FETCH_REDIRECT_STATUSES:237 await _drain_aiohttp_body_capped(238 response, _REDIRECT_RESPONSE_BODY_CAP_BYTES239 )240 if redirect_hops >= constants._MAX_WEB_FETCH_REDIRECTS:241 raise WebFetchEgressViolation(242 "web_fetch exceeded maximum redirects "243 f"({constants._MAX_WEB_FETCH_REDIRECTS})"244 )245 location = response.headers.get("location")246 if not location or not location.strip():247 raise WebFetchEgressViolation(248 "web_fetch redirect response missing Location header"249 )250 current_url = urljoin(str(response.url), location.strip())251 redirect_hops += 1252 continue253 response.raise_for_status()254 content_type = response.headers.get("content-type", "text/plain")255 final_url = str(response.url)256 encoding = response.get_encoding() or "utf-8"257 body_bytes = await _read_aiohttp_body_capped(258 response, constants._MAX_WEB_FETCH_RESPONSE_BYTES259 )260 finally:261 await connector.close()262 263 break264 265 text = body_bytes.decode(encoding, errors="replace")266 title = final_url267 data = text268 if "html" in content_type.lower():269 parser = HTMLTextParser()270 parser.feed(text)271 title = parser.title or final_url272 data = "\n".join(parser.text_parts)273 return {274 "url": final_url,275 "title": title,276 "media_type": "text/plain",277 "data": data[:_MAX_FETCH_CHARS],278 }279 