288 lines
12 KiB
Python
288 lines
12 KiB
Python
"""Explicit-manifest HTTP Evidence source with SSRF-safe bounded acquisition."""
|
|
|
|
import hashlib
|
|
import ipaddress
|
|
import socket
|
|
from collections import OrderedDict
|
|
from datetime import UTC, datetime
|
|
from email.utils import parsedate_to_datetime
|
|
from urllib.parse import urljoin, urlsplit
|
|
|
|
import requests
|
|
|
|
from tht.evidence.contracts import (
|
|
AcquiredDocument,
|
|
EvidenceSourceError,
|
|
EvidenceSourceErrorCategory,
|
|
SourceObject,
|
|
canonical_provenance_uri,
|
|
)
|
|
|
|
|
|
class HttpManifestEvidenceSource:
|
|
def __init__(
|
|
self,
|
|
urls: list[str] | tuple[str, ...],
|
|
*,
|
|
connect_timeout: float = 5,
|
|
read_timeout: float = 30,
|
|
max_bytes: int = 10 * 1024 * 1024,
|
|
max_redirects: int = 5,
|
|
allow_private_hosts: bool = False,
|
|
max_cache_bytes: int = 64 * 1024 * 1024,
|
|
) -> None:
|
|
if not urls:
|
|
raise ValueError("HTTP evidence manifest must contain at least one URL")
|
|
if (
|
|
connect_timeout <= 0
|
|
or read_timeout <= 0
|
|
or max_bytes < 1
|
|
or max_redirects < 0
|
|
or max_cache_bytes < 1
|
|
):
|
|
raise ValueError("HTTP evidence limits must be positive")
|
|
self._transport_by_uri: dict[str, str] = {}
|
|
for url in urls:
|
|
self._validate_url_shape(url)
|
|
provenance = canonical_provenance_uri(url)
|
|
if provenance in self._transport_by_uri:
|
|
raise ValueError("HTTP evidence manifest contains duplicate canonical provenance")
|
|
self._transport_by_uri[provenance] = url
|
|
self.connect_timeout = connect_timeout
|
|
self.read_timeout = read_timeout
|
|
self.max_bytes = max_bytes
|
|
self.max_redirects = max_redirects
|
|
self.allow_private_hosts = allow_private_hosts
|
|
self.max_cache_bytes = max_cache_bytes
|
|
self._session = requests.Session()
|
|
self._session.trust_env = False
|
|
self._cache: OrderedDict[str, AcquiredDocument] = OrderedDict()
|
|
# provenance -> (exact final effective URL, ETag, Last-Modified)
|
|
self._validators: dict[str, tuple[str, str | None, str | None]] = {}
|
|
self._cache_bytes = 0
|
|
|
|
def __repr__(self) -> str:
|
|
return f"HttpManifestEvidenceSource(objects={len(self._transport_by_uri)})"
|
|
|
|
@staticmethod
|
|
def _safe_error(operation: str, *, transient: bool = False, **details):
|
|
return EvidenceSourceError(
|
|
"HTTP source operation failed",
|
|
category=(
|
|
EvidenceSourceErrorCategory.TRANSIENT
|
|
if transient
|
|
else EvidenceSourceErrorCategory.PERMANENT
|
|
),
|
|
details={"operation": operation, **details},
|
|
)
|
|
|
|
@staticmethod
|
|
def _validate_url_shape(url: str) -> None:
|
|
parsed = urlsplit(url)
|
|
if parsed.scheme not in {"http", "https"} or not parsed.hostname:
|
|
raise ValueError("HTTP evidence URLs must use http or https")
|
|
if parsed.username is not None or parsed.password is not None:
|
|
raise ValueError("HTTP evidence URLs must not contain userinfo credentials")
|
|
|
|
@staticmethod
|
|
def _source_id(uri: str) -> str:
|
|
return f"http:{hashlib.sha256(uri.encode()).hexdigest()}"
|
|
|
|
@staticmethod
|
|
def _normalized_ip(value: str) -> ipaddress.IPv4Address | ipaddress.IPv6Address:
|
|
address = ipaddress.ip_address(value.split("%", 1)[0])
|
|
if isinstance(address, ipaddress.IPv6Address) and address.ipv4_mapped:
|
|
return address.ipv4_mapped
|
|
return address
|
|
|
|
def _resolve_allowed(self, url: str) -> set[ipaddress.IPv4Address | ipaddress.IPv6Address]:
|
|
try:
|
|
self._validate_url_shape(url)
|
|
except ValueError as error:
|
|
raise self._safe_error("url_validation") from error
|
|
if self.allow_private_hosts:
|
|
return set()
|
|
parsed = urlsplit(url)
|
|
port = parsed.port or (443 if parsed.scheme == "https" else 80)
|
|
try:
|
|
rows = socket.getaddrinfo(parsed.hostname, port, type=socket.SOCK_STREAM)
|
|
addresses = {self._normalized_ip(row[4][0]) for row in rows}
|
|
except (OSError, ValueError) as error:
|
|
raise self._safe_error("resolution", transient=True) from error
|
|
if not addresses:
|
|
raise self._safe_error("resolution", transient=True)
|
|
# Reject the entire answer set if any address is private/reserved. Choosing only a public
|
|
# member would leave DNS ordering as a policy bypass.
|
|
if any(not address.is_global for address in addresses):
|
|
raise self._safe_error("network_policy")
|
|
return addresses
|
|
|
|
def _verify_peer(
|
|
self,
|
|
response,
|
|
allowed: set[ipaddress.IPv4Address | ipaddress.IPv6Address],
|
|
) -> None:
|
|
if self.allow_private_hosts:
|
|
return
|
|
try:
|
|
connection = response.raw._connection
|
|
peer = self._normalized_ip(connection.sock.getpeername()[0])
|
|
except (AttributeError, OSError, TypeError, ValueError) as error:
|
|
raise self._safe_error("peer_validation", transient=True) from error
|
|
if not peer.is_global or peer not in allowed:
|
|
raise self._safe_error("network_policy")
|
|
|
|
@staticmethod
|
|
def _status_category(status: int) -> EvidenceSourceErrorCategory:
|
|
if status in {408, 425, 429} or 500 <= status <= 599:
|
|
return EvidenceSourceErrorCategory.TRANSIENT
|
|
return EvidenceSourceErrorCategory.PERMANENT
|
|
|
|
def _conditional_headers(self, provenance: str, request_url: str) -> dict[str, str]:
|
|
cached = self._cache.get(self._source_id(provenance))
|
|
validators = self._validators.get(provenance)
|
|
if cached is None or validators is None:
|
|
return {}
|
|
final_url, etag, last_modified = validators
|
|
if request_url != final_url:
|
|
return {}
|
|
headers = {}
|
|
if etag:
|
|
headers["If-None-Match"] = etag
|
|
if last_modified:
|
|
headers["If-Modified-Since"] = last_modified
|
|
return headers
|
|
|
|
def _remember(
|
|
self,
|
|
provenance: str,
|
|
final_url: str,
|
|
document: AcquiredDocument,
|
|
validators: tuple[str | None, str | None],
|
|
) -> None:
|
|
source_id = document.source.source_id
|
|
old = self._cache.pop(source_id, None)
|
|
if old is not None:
|
|
self._cache_bytes -= len(old.content)
|
|
self._cache[source_id] = document
|
|
self._cache_bytes += len(document.content)
|
|
self._validators[provenance] = (final_url, *validators)
|
|
while self._cache and self._cache_bytes > self.max_cache_bytes:
|
|
evicted_id, evicted = self._cache.popitem(last=False)
|
|
self._cache_bytes -= len(evicted.content)
|
|
for uri in tuple(self._validators):
|
|
if self._source_id(uri) == evicted_id:
|
|
del self._validators[uri]
|
|
|
|
def _download(self, transport_url: str, provenance: str) -> AcquiredDocument:
|
|
current = transport_url
|
|
try:
|
|
for redirect_count in range(self.max_redirects + 1):
|
|
headers = self._conditional_headers(provenance, current)
|
|
allowed = self._resolve_allowed(current)
|
|
response = None
|
|
try:
|
|
response = self._session.get(
|
|
current,
|
|
headers=headers,
|
|
stream=True,
|
|
allow_redirects=False,
|
|
timeout=(self.connect_timeout, self.read_timeout),
|
|
)
|
|
self._verify_peer(response, allowed)
|
|
if response.is_redirect:
|
|
if redirect_count == self.max_redirects:
|
|
raise self._safe_error("redirect")
|
|
destination = urljoin(current, response.headers.get("Location", ""))
|
|
try:
|
|
self._validate_url_shape(destination)
|
|
except ValueError as error:
|
|
raise self._safe_error("redirect") from error
|
|
# The next iteration binds validators to the exact destination URL.
|
|
current = destination
|
|
continue
|
|
if response.status_code == 304:
|
|
cached = self._cache.get(self._source_id(provenance))
|
|
binding = self._validators.get(provenance)
|
|
if (
|
|
cached is None
|
|
or not headers
|
|
or binding is None
|
|
or binding[0] != current
|
|
):
|
|
raise self._safe_error("conditional_response")
|
|
self._cache.move_to_end(cached.source.source_id)
|
|
return cached
|
|
if not 200 <= response.status_code <= 299:
|
|
raise EvidenceSourceError(
|
|
"HTTP status failure",
|
|
category=self._status_category(response.status_code),
|
|
details={"operation": "download", "status": response.status_code},
|
|
)
|
|
length = response.headers.get("Content-Length")
|
|
if length is not None and int(length) > self.max_bytes:
|
|
raise self._safe_error("download", limit_bytes=self.max_bytes)
|
|
content = bytearray()
|
|
for chunk in response.iter_content(
|
|
chunk_size=min(64 * 1024, self.max_bytes + 1)
|
|
):
|
|
content.extend(chunk)
|
|
if len(content) > self.max_bytes:
|
|
raise self._safe_error("download", limit_bytes=self.max_bytes)
|
|
etag = response.headers.get("ETag")
|
|
last_modified = response.headers.get("Last-Modified")
|
|
media_type = (
|
|
response.headers.get("Content-Type", "").split(";", 1)[0] or None
|
|
)
|
|
break
|
|
finally:
|
|
if response is not None:
|
|
response.close()
|
|
except EvidenceSourceError:
|
|
raise
|
|
except (requests.Timeout, requests.ConnectionError, TimeoutError) as error:
|
|
raise self._safe_error("download", transient=True) from error
|
|
except requests.RequestException as error:
|
|
raise self._safe_error("download", transient=True) from error
|
|
except (OSError, ValueError) as error:
|
|
raise self._safe_error("download") from error
|
|
|
|
modified_at = None
|
|
if etag:
|
|
fingerprint = f"etag:{hashlib.sha256(etag.encode()).hexdigest()}"
|
|
elif last_modified:
|
|
try:
|
|
modified_at = parsedate_to_datetime(last_modified).astimezone(UTC)
|
|
fingerprint = f"last-modified:{int(modified_at.timestamp())}"
|
|
except (TypeError, ValueError, OverflowError):
|
|
fingerprint = f"sha256:{hashlib.sha256(content).hexdigest()}"
|
|
else:
|
|
fingerprint = f"sha256:{hashlib.sha256(content).hexdigest()}"
|
|
item = SourceObject(
|
|
source_id=self._source_id(provenance),
|
|
uri=provenance,
|
|
fingerprint=fingerprint,
|
|
modified_at=modified_at,
|
|
)
|
|
document = AcquiredDocument(
|
|
source=item,
|
|
content=bytes(content),
|
|
media_type=media_type,
|
|
acquired_at=datetime.now(UTC),
|
|
)
|
|
self._remember(provenance, current, document, (etag, last_modified))
|
|
return document
|
|
|
|
def discover(self):
|
|
for provenance in sorted(self._transport_by_uri):
|
|
yield self._download(self._transport_by_uri[provenance], provenance).source
|
|
|
|
def acquire(self, item: SourceObject) -> AcquiredDocument:
|
|
transport = self._transport_by_uri.get(item.uri)
|
|
if transport is None or item.source_id != self._source_id(item.uri):
|
|
raise self._safe_error("acquire")
|
|
document = self._download(transport, item.uri)
|
|
if document.source.fingerprint != item.fingerprint:
|
|
raise self._safe_error("acquire")
|
|
return document
|