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():