diff --git a/s16code/tools.py b/s16code/tools.py index f70af95..f302b6c 100644 --- a/s16code/tools.py +++ b/s16code/tools.py @@ -6,16 +6,18 @@ import csv import hashlib import html +import ipaddress import math import operator import os import re import shutil +import socket import sqlite3 from datetime import date, datetime, timedelta from html.parser import HTMLParser from pathlib import Path -from urllib.parse import parse_qs, unquote, urlparse +from urllib.parse import parse_qs, unquote, urljoin, urlparse from zoneinfo import ZoneInfo, ZoneInfoNotFoundError import httpx @@ -244,21 +246,64 @@ def date_shift(value: str, days: int) -> dict: "weekday": shifted.strftime("%A")} -async def fetch_url(url: str, *, max_chars: int = 60_000) -> dict: +def assert_fetchable_url(url: str) -> None: + """Refuse anything that is not a public http(s) target. + + Scheme checks alone are not enough: ``http://127.0.0.1/`` and the cloud + metadata link-local address are valid http URLs and would otherwise let a + fetched page (or a redirect) read the host the agent is running on. + """ parsed = urlparse(url) if parsed.scheme not in {"http", "https"}: raise ValueError("fetch_url permits only http(s)") - async with httpx.AsyncClient(timeout=30, follow_redirects=True, + host = parsed.hostname + if not host: + raise ValueError("fetch_url URL has no host") + try: + addresses = [ipaddress.ip_address(host)] + except ValueError: + try: + infos = socket.getaddrinfo(host, parsed.port or (443 if parsed.scheme == "https" else 80), + type=socket.SOCK_STREAM) + except socket.gaierror as error: + raise ValueError(f"fetch_url could not resolve host: {host}") from error + addresses = [ipaddress.ip_address(info[4][0]) for info in infos] + if not addresses: + raise ValueError(f"fetch_url could not resolve host: {host}") + for address in addresses: + if isinstance(address, ipaddress.IPv6Address) and address.ipv4_mapped is not None: + address = address.ipv4_mapped + if (address.is_loopback or address.is_private or address.is_link_local + or address.is_multicast or address.is_reserved or address.is_unspecified): + raise ValueError("fetch_url refuses private or loopback addresses") + + +async def fetch_url(url: str, *, max_chars: int = 60_000) -> dict: + current = url + async with httpx.AsyncClient(timeout=30, follow_redirects=False, headers={"User-Agent": "GLC-S16/0.3 educational-agent"}) as client: - response = await client.get(url) - response.raise_for_status() - content_type = response.headers.get("content-type", "") - if "html" in content_type: - parser = _TextExtractor(); parser.feed(response.text); text = parser.text() - else: - text = response.text - return {"url": str(response.url), "status": response.status_code, - "content_type": content_type, "text": text[:max_chars], "truncated": len(text) > max_chars} + response = None + for _ in range(5): + assert_fetchable_url(current) + response = await client.get(current) + if response.is_redirect: + location = response.headers.get("location") + if not location: + raise ValueError("fetch_url redirect was missing Location") + current = urljoin(str(response.url), location) + continue + response.raise_for_status() + break + else: + raise ValueError("fetch_url followed too many redirects") + assert response is not None + content_type = response.headers.get("content-type", "") + if "html" in content_type: + parser = _TextExtractor(); parser.feed(response.text); text = parser.text() + else: + text = response.text + return {"url": str(response.url), "status": response.status_code, + "content_type": content_type, "text": text[:max_chars], "truncated": len(text) > max_chars} class _DDGParser(HTMLParser): diff --git a/tests/test_general_tools.py b/tests/test_general_tools.py index 8b08223..1412da9 100644 --- a/tests/test_general_tools.py +++ b/tests/test_general_tools.py @@ -24,6 +24,25 @@ def test_calculate_supports_useful_arithmetic_without_code_execution(): calculate("2 ** 100") +def test_fetch_url_refuses_private_and_loopback_targets(): + from s16code.tools import assert_fetchable_url + + assert_fetchable_url("https://1.1.1.1/") + with pytest.raises(ValueError, match="http"): + assert_fetchable_url("file:///etc/passwd") + for hostile in ( + "http://127.0.0.1/", + "http://127.0.0.1:8111/v1/budget", + "http://localhost/", + "http://10.0.0.1/", + "http://192.168.1.1/", + "http://169.254.169.254/latest/meta-data/", + "http://[::1]/", + ): + with pytest.raises(ValueError, match="private or loopback"): + assert_fetchable_url(hostile) + + def test_write_and_hash_stay_inside_sandbox(monkeypatch, tmp_path): monkeypatch.setenv("S16_SANDBOX_ROOT", str(tmp_path)) written = write_text_file("reports/result.txt", "verified output")