fix(evidence): close S3 and smoke safety gaps

This commit is contained in:
2026-07-12 06:09:15 +02:00
parent efcb0deb31
commit e6d44ba082
5 changed files with 148 additions and 36 deletions
+60 -4
View File
@@ -12,8 +12,12 @@ class Body:
class Client:
def __init__(self): self.body = Body(b"hello")
def __init__(self): self.body, self.list_calls = Body(b"hello"), 0
def get_paginator(self, name): return self
def list_objects_v2(self, **kwargs):
self.list_calls += 1
return {"Contents": [{"Key": "clinical/a.md", "ETag": '"abc"',
"Size": 5, "LastModified": datetime(2026, 1, 1, tzinfo=UTC)}]}
def paginate(self, **kwargs):
yield {"Contents": [{"Key": "clinical/a.md", "ETag": '"abc"',
"Size": 5, "LastModified": datetime(2026, 1, 1, tzinfo=UTC)}]}
@@ -60,6 +64,13 @@ def test_s3_rejects_invalid_bucket_names(bucket):
S3EvidenceSource(bucket=bucket, client=Client())
@pytest.mark.parametrize("bucket", ["127.0.0.1", "192.168.1.1"])
def test_s3_rejects_ip_shaped_bucket(bucket):
from tht.adapters.evidence.s3 import S3EvidenceSource
with pytest.raises(ValueError, match="bucket"):
S3EvidenceSource(bucket=bucket, client=Client())
def test_s3_rejects_endpoint_query_path_fragment_and_untrusted_custom_host():
from tht.adapters.evidence.s3 import S3EvidenceSource
for endpoint in ("https://s3.example.test/path", "https://s3.example.test/?x=1",
@@ -74,9 +85,37 @@ def test_s3_rejects_endpoint_query_path_fragment_and_untrusted_custom_host():
def test_s3_rejects_out_of_prefix_key_and_missing_validator():
from tht.adapters.evidence.s3 import S3EvidenceSource
client = Client()
client.paginate = lambda **kwargs: iter([{"Contents": [{"Key": "other/a.md", "ETag": '"x"'}]}])
client.list_objects_v2 = lambda **kwargs: {"Contents": [{"Key": "other/a.md", "ETag": '"x"'}]}
with pytest.raises(EvidenceSourceError):
list(S3EvidenceSource(bucket="evidence", prefix="clinical/", client=client).discover())
client.list_objects_v2 = lambda **kwargs: {"Contents": [{"Key": "clinical/a.md"}]}
with pytest.raises(EvidenceSourceError):
list(S3EvidenceSource(bucket="evidence", prefix="clinical/", client=client).discover())
def test_s3_rejects_leading_slash_prefix_empty_and_control_keys():
from tht.adapters.evidence.s3 import S3EvidenceSource
with pytest.raises(ValueError, match="prefix"):
S3EvidenceSource(bucket="evidence", prefix="/clinical", client=Client())
for key in ("", "clinical/a\x00.md", "clinical/a\x7f.md"):
client = Client()
client.list_objects_v2 = lambda **kwargs: {"Contents": [{"Key": key, "ETag": '"x"'}]}
with pytest.raises(EvidenceSourceError):
list(S3EvidenceSource(bucket="evidence", prefix="clinical/", client=client).discover())
def test_s3_hard_page_limit_never_requests_page_max_plus_one():
from tht.adapters.evidence.s3 import S3EvidenceSource
client = Client()
def listing(**kwargs):
client.list_calls += 1
return {"Contents": [{"Key": f"clinical/{client.list_calls}.md", "ETag": '"x"'}],
"IsTruncated": True, "NextContinuationToken": str(client.list_calls)}
client.list_objects_v2 = listing
with pytest.raises(EvidenceSourceError):
list(S3EvidenceSource(bucket="evidence", prefix="clinical/", max_pages=2,
client=client).discover())
assert client.list_calls == 2
def test_s3_acquire_rejects_exact_etag_drift_and_closes_body():
@@ -89,9 +128,26 @@ def test_s3_acquire_rejects_exact_etag_drift_and_closes_body():
with pytest.raises(EvidenceSourceError):
source.acquire(item)
assert client.body.closed
client.paginate = lambda **kwargs: iter([{"Contents": [{"Key": "clinical/a.md"}]}])
def test_s3_acquire_rejects_forged_reconstructed_item_before_get():
from tht.adapters.evidence.s3 import S3EvidenceSource
client = Client()
source = S3EvidenceSource(bucket="evidence", client=client)
item = next(iter(source.discover()))
forged = item.model_copy(update={"fingerprint": "etag:" + "0" * 64})
client.get_object = lambda **kwargs: (_ for _ in ()).throw(AssertionError("called"))
with pytest.raises(EvidenceSourceError):
list(S3EvidenceSource(bucket="evidence", prefix="clinical/", client=client).discover())
source.acquire(forged)
@pytest.mark.parametrize("host", ["127.0.0.1", "10.0.0.1", "169.254.1.1", "0.0.0.0",
"[::1]", "[fe80::1]", "[::]"])
def test_s3_literal_non_global_endpoint_requires_private_opt_in(host):
from tht.adapters.evidence.s3 import S3EvidenceSource
with pytest.raises(ValueError, match="private"):
S3EvidenceSource(bucket="evidence", endpoint_url=f"https://{host}:9000",
trusted_endpoint=True, client=Client())
def test_s3_size_limit_closes_body():
+45 -26
View File
@@ -1,6 +1,7 @@
"""Bounded S3-compatible Evidence source using the supported boto3 client."""
import hashlib
import ipaddress
import re
from datetime import UTC, datetime
from urllib.parse import quote, urlsplit
@@ -18,10 +19,18 @@ class S3EvidenceSource:
allow_private_endpoint: bool = False, allow_insecure_endpoint: bool = False,
max_bytes: int = 10 * 1024 * 1024, max_objects: int = 10_000,
max_pages: int = 100, page_size: int = 1000, client=None) -> None:
if (not re.fullmatch(r"(?=.{3,63}$)(?!-)(?!.*\.\.)(?!.*\.-)(?!.*-\.)"
bucket_valid = re.fullmatch(r"(?=.{3,63}$)(?!-)(?!.*\.\.)(?!.*\.-)(?!.*-\.)"
r"[a-z0-9](?:[a-z0-9.-]*[a-z0-9])?", bucket)
try:
ipaddress.ip_address(bucket)
bucket_is_ip = True
except ValueError:
bucket_is_ip = False
if (not bucket_valid or bucket_is_ip
or any(value < 1 for value in (max_bytes, max_objects, max_pages, page_size))):
raise ValueError("S3 evidence limits and bucket must be non-empty and positive")
if prefix.startswith("/"):
raise ValueError("S3 prefix must not start with a slash")
if endpoint_url:
parsed = urlsplit(endpoint_url)
if parsed.username or parsed.password:
@@ -36,9 +45,13 @@ class S3EvidenceSource:
raise ValueError("S3 custom endpoint must be an origin root without query/fragment")
if not trusted_endpoint:
raise ValueError("S3 custom endpoint requires explicit trusted_endpoint opt-in")
if parsed.hostname in {"localhost", "127.0.0.1", "::1"} and not allow_private_endpoint:
try:
literal = ipaddress.ip_address(parsed.hostname)
except ValueError:
literal = None
if literal is not None and not literal.is_global and not allow_private_endpoint:
raise ValueError("S3 private endpoint requires explicit opt-in")
self.bucket, self.prefix = bucket, prefix.lstrip("/")
self.bucket, self.prefix = bucket, prefix
self.max_bytes, self.max_objects = max_bytes, max_objects
self.max_pages, self.page_size = max_pages, min(page_size, 1000)
if client is None:
@@ -53,7 +66,7 @@ class S3EvidenceSource:
aws_session_token=session_token, verify=True,
config=BotoConfig(s3={"addressing_style": "path"}))
self._client = client
self._items: dict[str, tuple[str, str | None, str | None]] = {}
self._items: dict[str, tuple[SourceObject, str]] = {}
@staticmethod
def _error(operation: str, transient: bool = False):
@@ -65,32 +78,42 @@ class S3EvidenceSource:
def discover(self):
count = pages = 0
try:
paginator = self._client.get_paginator("list_objects_v2")
for page in paginator.paginate(Bucket=self.bucket, Prefix=self.prefix,
PaginationConfig={"PageSize": self.page_size}):
token = None
for _ in range(self.max_pages):
params = {"Bucket": self.bucket, "Prefix": self.prefix,
"MaxKeys": self.page_size}
if token is not None:
params["ContinuationToken"] = token
page = self._client.list_objects_v2(**params)
pages += 1
if pages > self.max_pages:
raise self._error("list_limit")
for row in page.get("Contents", []):
count += 1
if count > self.max_objects:
raise self._error("object_limit")
key, etag = row.get("Key"), row.get("ETag")
if (not isinstance(key, str) or not key.startswith(self.prefix)
or len(key.encode()) > 1024 or any(ord(char) < 32 for char in key)):
if (not isinstance(key, str) or not key or not key.startswith(self.prefix)
or len(key.encode()) > 1024
or any(ord(char) < 32 or ord(char) == 127 for char in key)):
raise self._error("invalid_key")
if not isinstance(etag, str) or not etag or len(etag) > 1024:
raise self._error("missing_validator")
uri = f"s3://{self.bucket}/{quote(key, safe='/')}"
fingerprint = f"etag:{hashlib.sha256(etag.encode()).hexdigest()}"
source_id = "s3:" + hashlib.sha256(uri.encode()).hexdigest()
self._items[source_id] = (key, None, etag)
modified = row.get("LastModified")
if modified is not None:
modified = modified.astimezone(UTC)
yield SourceObject(source_id=source_id, uri=uri, fingerprint=fingerprint,
modified_at=modified,
metadata={"size": int(row.get("Size", 0))})
item = SourceObject(source_id=source_id, uri=uri, fingerprint=fingerprint,
modified_at=modified,
metadata={"size": int(row.get("Size", 0))})
self._items[source_id] = (item, etag)
yield item
if not page.get("IsTruncated"):
return
token = page.get("NextContinuationToken")
if not isinstance(token, str) or not token:
raise self._error("list_continuation")
raise self._error("list_limit")
except EvidenceSourceError:
raise
except Exception as exc:
@@ -98,23 +121,19 @@ class S3EvidenceSource:
def acquire(self, item: SourceObject) -> AcquiredDocument:
binding = self._items.get(item.source_id)
if binding is None:
if binding is None or item != binding[0]:
raise self._error("acquire")
key, version, _etag = binding
discovered, etag = binding
key = discovered.uri.split(f"s3://{self.bucket}/", 1)[1]
from urllib.parse import unquote
key = unquote(key)
kwargs = {"Bucket": self.bucket, "Key": key}
if version:
kwargs["VersionId"] = version
body = None
try:
response = self._client.get_object(**kwargs)
body = response["Body"]
if version:
if response.get("VersionId") != version:
raise self._error("version_changed")
else:
current = hashlib.sha256((response.get("ETag") or "").encode()).hexdigest()
if item.fingerprint != f"etag:{current}":
raise self._error("etag_changed")
if response.get("ETag") != etag:
raise self._error("etag_changed")
if int(response.get("ContentLength", 0)) > self.max_bytes:
raise self._error("download_limit")
content = body.read(self.max_bytes + 1)