Team Ai
Apppublic

Yash030/claude-code-proxy

sourceHugging Faceupdated 5mo agoView on Hugging Face
2likes
outbound.py279 linesDownload Raw Back to web_tools
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