fix(evidence): harden source acquisition

This commit is contained in:
2026-07-12 03:32:18 +02:00
parent ffd683c587
commit 81ff1810d1
8 changed files with 523 additions and 183 deletions
+99 -75
View File
@@ -1,8 +1,11 @@
"""Contained, deterministic filesystem Evidence source."""
"""Contained, race-safe filesystem Evidence source."""
import hashlib
import os
import stat
from datetime import UTC, datetime
from pathlib import Path
from pathlib import Path, PurePosixPath
from urllib.parse import unquote, urlsplit
from tht.ports.evidence import (
AcquiredDocument,
@@ -22,103 +25,124 @@ class FilesystemEvidenceSource:
) -> None:
if max_bytes < 1:
raise ValueError("max_bytes must be positive")
if not patterns or any(not pattern for pattern in patterns):
if not patterns or any(
not pattern or Path(pattern).is_absolute() or ".." in Path(pattern).parts
for pattern in patterns
):
raise ValueError("at least one non-empty discovery pattern is required")
try:
self.root = Path(root).expanduser().resolve(strict=True)
self._root_fd = os.open(
self.root,
os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW | os.O_CLOEXEC,
)
except OSError as error:
raise ValueError("filesystem evidence root is unavailable") from error
if not self.root.is_dir():
raise ValueError("filesystem evidence root must be a directory")
self.patterns = tuple(patterns)
self.max_bytes = max_bytes
def _contained(self, path: Path) -> Path:
try:
resolved = path.resolve(strict=True)
resolved.relative_to(self.root)
except (OSError, ValueError) as error:
raise EvidenceSourceError(
"unsafe filesystem object",
category=EvidenceSourceErrorCategory.PERMANENT,
details={"operation": "path_validation"},
) from error
if not resolved.is_file():
raise EvidenceSourceError(
"unsupported filesystem object",
category=EvidenceSourceErrorCategory.PERMANENT,
details={"operation": "path_validation"},
)
return resolved
def __del__(self):
root_fd = getattr(self, "_root_fd", None)
if root_fd is not None:
try:
os.close(root_fd)
except OSError:
pass
def _read(self, path: Path) -> bytes:
@staticmethod
def _safe_error(operation: str, *, transient: bool = False, **details):
return EvidenceSourceError(
"filesystem source operation failed",
category=(
EvidenceSourceErrorCategory.TRANSIENT
if transient
else EvidenceSourceErrorCategory.PERMANENT
),
details={"operation": operation, **details},
)
def _open_read(self, relative: PurePosixPath) -> tuple[bytes, os.stat_result]:
parts = relative.parts
if not parts or any(part in {"", ".", ".."} for part in parts):
raise self._safe_error("path_validation")
directory_fd = os.dup(self._root_fd)
file_fd = None
try:
if path.stat().st_size > self.max_bytes:
raise EvidenceSourceError(
"filesystem object exceeds configured limit",
category=EvidenceSourceErrorCategory.PERMANENT,
details={"operation": "read", "limit_bytes": self.max_bytes},
for component in parts[:-1]:
next_fd = os.open(
component,
os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW | os.O_CLOEXEC,
dir_fd=directory_fd,
)
with path.open("rb") as stream:
content = stream.read(self.max_bytes + 1)
os.close(directory_fd)
directory_fd = next_fd
file_fd = os.open(
parts[-1],
os.O_RDONLY | os.O_NOFOLLOW | os.O_CLOEXEC,
dir_fd=directory_fd,
)
file_stat = os.fstat(file_fd)
if not stat.S_ISREG(file_stat.st_mode):
raise self._safe_error("path_validation")
if file_stat.st_size > self.max_bytes:
raise self._safe_error("read", limit_bytes=self.max_bytes)
content = bytearray()
while len(content) <= self.max_bytes:
chunk = os.read(file_fd, min(64 * 1024, self.max_bytes + 1 - len(content)))
if not chunk:
break
content.extend(chunk)
if len(content) > self.max_bytes:
raise self._safe_error("read", limit_bytes=self.max_bytes)
return bytes(content), file_stat
except EvidenceSourceError:
raise
except OSError as error:
raise EvidenceSourceError(
"filesystem read failed",
category=EvidenceSourceErrorCategory.TRANSIENT,
details={"operation": "read"},
) from error
if len(content) > self.max_bytes:
raise EvidenceSourceError(
"filesystem object exceeds configured limit",
category=EvidenceSourceErrorCategory.PERMANENT,
details={"operation": "read", "limit_bytes": self.max_bytes},
)
return content
raise self._safe_error("open") from error
finally:
if file_fd is not None:
os.close(file_fd)
os.close(directory_fd)
def _item(self, path: Path, content: bytes) -> SourceObject:
relative = path.relative_to(self.root).as_posix()
digest = hashlib.sha256(content).hexdigest()
stable_id = hashlib.sha256(relative.encode()).hexdigest()
modified = datetime.fromtimestamp(path.stat().st_mtime, tz=UTC)
def _item(
self, relative: PurePosixPath, content: bytes, file_stat: os.stat_result
) -> SourceObject:
relative_text = relative.as_posix()
return SourceObject(
source_id=f"filesystem:{stable_id}",
uri=path.as_uri(),
fingerprint=f"sha256:{digest}",
modified_at=modified,
metadata={"relative_path": relative},
source_id=f"filesystem:{hashlib.sha256(relative_text.encode()).hexdigest()}",
uri=(self.root / relative_text).as_uri(),
fingerprint=f"sha256:{hashlib.sha256(content).hexdigest()}",
modified_at=datetime.fromtimestamp(file_stat.st_mtime, tz=UTC),
metadata={"relative_path": relative_text},
)
def discover(self):
candidates = {path for pattern in self.patterns for path in self.root.glob(pattern)}
for candidate in sorted(candidates, key=lambda path: path.as_posix()):
path = self._contained(candidate)
content = self._read(path)
yield self._item(path, content)
candidates = {
path.relative_to(self.root).as_posix()
for pattern in self.patterns
for path in self.root.glob(pattern)
}
for relative_text in sorted(candidates):
relative = PurePosixPath(relative_text)
content, file_stat = self._open_read(relative)
yield self._item(relative, content, file_stat)
def acquire(self, item: SourceObject) -> AcquiredDocument:
if not item.uri.startswith("file:"):
raise EvidenceSourceError(
"object does not belong to filesystem source",
category=EvidenceSourceErrorCategory.PERMANENT,
details={"operation": "acquire"},
)
from urllib.parse import unquote, urlsplit
parsed = urlsplit(item.uri)
path = self._contained(Path(unquote(parsed.path)))
content = self._read(path)
expected = self._item(path, content)
if item.source_id != expected.source_id:
raise EvidenceSourceError(
"object does not belong to filesystem source",
category=EvidenceSourceErrorCategory.PERMANENT,
details={"operation": "acquire"},
)
if parsed.scheme != "file" or parsed.netloc or parsed.query or parsed.fragment:
raise self._safe_error("acquire")
try:
relative = Path(unquote(parsed.path)).relative_to(self.root)
except ValueError as error:
raise self._safe_error("acquire") from error
pure_relative = PurePosixPath(relative.as_posix())
content, file_stat = self._open_read(pure_relative)
expected = self._item(pure_relative, content, file_stat)
if item.source_id != expected.source_id or item.fingerprint != expected.fingerprint:
raise self._safe_error("acquire")
return AcquiredDocument(
source=expected,
content=content,
media_type="text/markdown" if path.suffix.lower() == ".md" else None,
media_type="text/markdown" if relative.suffix.lower() == ".md" else None,
acquired_at=datetime.now(UTC),
)
+191 -103
View File
@@ -1,8 +1,9 @@
"""Explicit-manifest HTTP Evidence source with bounded streaming reads."""
"""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
@@ -27,18 +28,22 @@ class HttpManifestEvidenceSource:
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:
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:
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")
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")
@@ -47,39 +52,84 @@ class HttpManifestEvidenceSource:
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._cache: dict[str, AcquiredDocument] = {}
self._session.trust_env = False
self._cache: OrderedDict[str, AcquiredDocument] = OrderedDict()
self._validators: dict[str, tuple[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 _reject_private_redirect(url: str) -> None:
parsed = urlsplit(url)
if parsed.scheme not in {"http", "https"} or not parsed.hostname:
raise EvidenceSourceError(
"redirect uses unsupported destination",
category=EvidenceSourceErrorCategory.PERMANENT,
details={"operation": "redirect"},
)
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:
addresses = {row[4][0] for row in socket.getaddrinfo(parsed.hostname, parsed.port)}
except OSError as error:
raise EvidenceSourceError(
"redirect destination resolution failed",
category=EvidenceSourceErrorCategory.TRANSIENT,
details={"operation": "redirect_resolution"},
) from error
if any(not ipaddress.ip_address(address).is_global for address in addresses):
raise EvidenceSourceError(
"redirect to private destination is forbidden",
category=EvidenceSourceErrorCategory.PERMANENT,
details={"operation": "redirect"},
)
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:
@@ -87,78 +137,121 @@ class HttpManifestEvidenceSource:
return EvidenceSourceErrorCategory.TRANSIENT
return EvidenceSourceErrorCategory.PERMANENT
def _conditional_headers(self, provenance: 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 {}
etag, last_modified = validators
headers = {}
if etag:
headers["If-None-Match"] = etag
if last_modified:
headers["If-Modified-Since"] = last_modified
return headers
def _remember(
self,
provenance: 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] = 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
headers = self._conditional_headers(provenance)
try:
for redirect_count in range(self.max_redirects + 1):
response = self._session.get(
current,
stream=True,
allow_redirects=False,
timeout=(self.connect_timeout, self.read_timeout),
)
if response.is_redirect:
response.close()
if redirect_count == self.max_redirects:
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
current_origin = urlsplit(current)
destination_origin = urlsplit(destination)
if (
current_origin.scheme,
current_origin.hostname,
current_origin.port,
) != (
destination_origin.scheme,
destination_origin.hostname,
destination_origin.port,
):
headers = {}
# Resolve every hop independently. Validators are never forwarded across
# origins, where even an opaque ETag would become cross-origin state.
current = destination
continue
if response.status_code == 304:
cached = self._cache.get(self._source_id(provenance))
if cached is None:
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(
"too many redirects",
category=EvidenceSourceErrorCategory.PERMANENT,
details={"operation": "redirect"},
"HTTP status failure",
category=self._status_category(response.status_code),
details={"operation": "download", "status": response.status_code},
)
destination = urljoin(current, response.headers.get("Location", ""))
self._reject_private_redirect(destination)
current = destination
continue
if not 200 <= response.status_code <= 299:
status = response.status_code
response.close()
raise EvidenceSourceError(
"HTTP status failure",
category=self._status_category(status),
details={"operation": "download", "status": status},
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
)
length = response.headers.get("Content-Length")
if length is not None and int(length) > self.max_bytes:
response.close()
raise EvidenceSourceError(
"HTTP object exceeds configured limit",
category=EvidenceSourceErrorCategory.PERMANENT,
details={"operation": "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:
break
finally:
if response is not None:
response.close()
raise EvidenceSourceError(
"HTTP object exceeds configured limit",
category=EvidenceSourceErrorCategory.PERMANENT,
details={"operation": "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
response.close()
break
except EvidenceSourceError:
raise
except (requests.Timeout, requests.ConnectionError) as error:
raise EvidenceSourceError(
"HTTP transport unavailable",
category=EvidenceSourceErrorCategory.TRANSIENT,
details={"operation": "download"},
) from error
except (requests.RequestException, ValueError) as error:
raise EvidenceSourceError(
"HTTP acquisition failed",
category=EvidenceSourceErrorCategory.PERMANENT,
details={"operation": "download"},
) from error
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:
# Entity tags are commonly quoted; the port's stable-value grammar is deliberately
# narrower, so preserve the opaque validator through a deterministic digest.
fingerprint = f"etag:{hashlib.sha256(etag.encode()).hexdigest()}"
elif last_modified:
try:
@@ -174,29 +267,24 @@ class HttpManifestEvidenceSource:
fingerprint=fingerprint,
modified_at=modified_at,
)
return AcquiredDocument(
document = AcquiredDocument(
source=item,
content=bytes(content),
media_type=media_type,
acquired_at=datetime.now(UTC),
)
self._remember(provenance, document, (etag, last_modified))
return document
def discover(self):
for provenance in sorted(self._transport_by_uri):
document = self._download(self._transport_by_uri[provenance], provenance)
self._cache[document.source.source_id] = document
yield document.source
yield self._download(self._transport_by_uri[provenance], provenance).source
def acquire(self, item: SourceObject) -> AcquiredDocument:
expected_id = self._source_id(item.uri)
transport = self._transport_by_uri.get(item.uri)
if transport is None or item.source_id != expected_id:
raise EvidenceSourceError(
"object does not belong to HTTP source",
category=EvidenceSourceErrorCategory.PERMANENT,
details={"operation": "acquire"},
)
cached = self._cache.get(item.source_id)
if cached is not None and cached.source.fingerprint == item.fingerprint:
return cached
return self._download(transport, 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
+2
View File
@@ -113,6 +113,8 @@ def build_evidence_sources(cfg: Config):
read_timeout=resource.read_timeout,
max_bytes=resource.max_bytes,
max_redirects=resource.max_redirects,
allow_private_hosts=resource.allow_private_hosts,
max_cache_bytes=resource.max_cache_bytes,
)
)
case other: # pragma: no cover - Pydantic rejects unsupported discriminators.