fix(preprocess): harden S3 and real Compose jobs

This commit is contained in:
2026-07-12 06:01:06 +02:00
parent 950f88f23e
commit c6966f3d15
10 changed files with 222 additions and 43 deletions
+50 -5
View File
@@ -15,12 +15,12 @@ class Client:
def __init__(self): self.body = Body(b"hello")
def get_paginator(self, name): return self
def paginate(self, **kwargs):
yield {"Contents": [{"Key": "clinical/a.md", "ETag": '"abc"', "VersionId": "v1",
yield {"Contents": [{"Key": "clinical/a.md", "ETag": '"abc"',
"Size": 5, "LastModified": datetime(2026, 1, 1, tzinfo=UTC)}]}
def get_object(self, **kwargs):
assert kwargs == {"Bucket": "evidence", "Key": "clinical/a.md", "VersionId": "v1"}
assert kwargs == {"Bucket": "evidence", "Key": "clinical/a.md"}
return {"Body": self.body, "ContentLength": 5, "ContentType": "text/markdown",
"ETag": '"abc"', "VersionId": "v1"}
"ETag": '"abc"'}
def test_s3_canonical_uri_version_fingerprint_and_closed_body():
@@ -29,7 +29,7 @@ def test_s3_canonical_uri_version_fingerprint_and_closed_body():
source = S3EvidenceSource(bucket="evidence", prefix="clinical/", client=client)
item = next(iter(source.discover()))
assert item.uri == "s3://evidence/clinical/a.md"
assert item.fingerprint == "s3-version:v1"
assert item.fingerprint.startswith("etag:")
assert source.acquire(item).content == b"hello"
assert client.body.closed
@@ -43,10 +43,55 @@ def test_s3_etag_fallback_and_bounds():
def test_s3_rejects_private_or_insecure_endpoint_without_explicit_opt_in():
from tht.adapters.evidence.s3 import S3EvidenceSource
with pytest.raises(ValueError, match="private"):
with pytest.raises(ValueError, match="trusted"):
S3EvidenceSource(bucket="evidence", endpoint_url="https://127.0.0.1:9000", client=Client())
with pytest.raises(ValueError, match="HTTPS"):
S3EvidenceSource(bucket="evidence", endpoint_url="http://s3.example.test", client=Client())
source = S3EvidenceSource(bucket="evidence", endpoint_url="http://127.0.0.1:9000",
trusted_endpoint=True, allow_private_endpoint=True,
allow_insecure_endpoint=True, client=Client())
assert source is not None
@pytest.mark.parametrize("bucket", ["UPPER", "bad_bucket", "-start", "end-", "a..b"])
def test_s3_rejects_invalid_bucket_names(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",
"https://s3.example.test/#x"):
with pytest.raises(ValueError, match="root"):
S3EvidenceSource(bucket="evidence", endpoint_url=endpoint,
trusted_endpoint=True, client=Client())
with pytest.raises(ValueError, match="trusted"):
S3EvidenceSource(bucket="evidence", endpoint_url="https://s3.example.test", client=Client())
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"'}]}])
with pytest.raises(EvidenceSourceError):
list(S3EvidenceSource(bucket="evidence", prefix="clinical/", client=client).discover())
def test_s3_acquire_rejects_exact_etag_drift_and_closes_body():
from tht.adapters.evidence.s3 import S3EvidenceSource
client = Client()
source = S3EvidenceSource(bucket="evidence", client=client)
item = next(iter(source.discover()))
client.get_object = lambda **kwargs: {"Body": client.body, "ContentLength": 5,
"ETag": '"changed"'}
with pytest.raises(EvidenceSourceError):
source.acquire(item)
assert client.body.closed
client.paginate = lambda **kwargs: iter([{"Contents": [{"Key": "clinical/a.md"}]}])
with pytest.raises(EvidenceSourceError):
list(S3EvidenceSource(bucket="evidence", prefix="clinical/", client=client).discover())
def test_s3_size_limit_closes_body():
+25 -18
View File
@@ -1,8 +1,7 @@
"""Bounded S3-compatible Evidence source using the supported boto3 client."""
import hashlib
import ipaddress
import socket
import re
from datetime import UTC, datetime
from urllib.parse import quote, urlsplit
@@ -15,40 +14,44 @@ class S3EvidenceSource:
def __init__(self, *, bucket: str, prefix: str = "", endpoint_url: str | None = None,
region: str | None = None, access_key: str | None = None,
secret_key: str | None = None, session_token: str | None = None,
trusted_endpoint: bool = False,
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 bucket or any(value < 1 for value in (max_bytes, max_objects, max_pages, page_size)):
if (not re.fullmatch(r"(?=.{3,63}$)(?!-)(?!.*\.\.)(?!.*\.-)(?!.*-\.)"
r"[a-z0-9](?:[a-z0-9.-]*[a-z0-9])?", bucket)
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 endpoint_url:
parsed = urlsplit(endpoint_url)
if parsed.username or parsed.password:
raise ValueError("S3 endpoint must not contain credentials")
if parsed.scheme != "https" and not allow_insecure_endpoint:
if parsed.scheme not in {"http", "https"}:
raise ValueError("S3 endpoint scheme must be exactly https or explicitly allowed http")
if parsed.scheme == "http" and not allow_insecure_endpoint:
raise ValueError("S3 endpoint must use HTTPS unless explicitly allowed")
if not parsed.hostname:
raise ValueError("S3 endpoint must include a hostname")
if not allow_private_endpoint:
try:
addresses = {ipaddress.ip_address(row[4][0].split("%", 1)[0]) for row in
socket.getaddrinfo(parsed.hostname, parsed.port or 443,
type=socket.SOCK_STREAM)}
except (OSError, ValueError) as exc:
raise ValueError("S3 endpoint resolution failed") from exc
if not addresses or any(not address.is_global for address in addresses):
raise ValueError("S3 private endpoint requires explicit opt-in")
if parsed.path not in {"", "/"} or parsed.query or parsed.fragment:
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:
raise ValueError("S3 private endpoint requires explicit opt-in")
self.bucket, self.prefix = bucket, prefix.lstrip("/")
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:
try:
import boto3
from botocore.config import Config as BotoConfig
except ImportError as exc: # pragma: no cover - deployment optional dependency
raise RuntimeError("Install tht[s3] to use S3 Evidence") from exc
client = boto3.client("s3", endpoint_url=endpoint_url, region_name=region,
aws_access_key_id=access_key,
aws_secret_access_key=secret_key,
aws_session_token=session_token, verify=True)
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]] = {}
@@ -72,12 +75,16 @@ class S3EvidenceSource:
count += 1
if count > self.max_objects:
raise self._error("object_limit")
key, version, etag = row["Key"], row.get("VersionId"), row.get("ETag")
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)):
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='/')}"
stable = version or hashlib.sha256((etag or "").encode()).hexdigest()
fingerprint = f"s3-version:{stable}" if version else f"etag:{stable}"
fingerprint = f"etag:{hashlib.sha256(etag.encode()).hexdigest()}"
source_id = "s3:" + hashlib.sha256(uri.encode()).hexdigest()
self._items[source_id] = (key, version, etag)
self._items[source_id] = (key, None, etag)
modified = row.get("LastModified")
if modified is not None:
modified = modified.astimezone(UTC)
+1
View File
@@ -127,6 +127,7 @@ def build_evidence_sources(cfg: Config):
endpoint_url=resource.endpoint_url, region=resource.region,
access_key=secret(resource.access_key), secret_key=secret(resource.secret_key),
session_token=secret(resource.session_token),
trusted_endpoint=resource.trusted_endpoint,
allow_private_endpoint=resource.allow_private_endpoint,
allow_insecure_endpoint=resource.allow_insecure_endpoint,
max_bytes=resource.max_bytes, max_objects=resource.max_objects,
+1
View File
@@ -204,6 +204,7 @@ class S3EvidenceSourceConfig(BaseModel):
access_key: SecretStr | None = None
secret_key: SecretStr | None = None
session_token: SecretStr | None = None
trusted_endpoint: bool = False
allow_private_endpoint: bool = False
allow_insecure_endpoint: bool = False
max_bytes: int = Field(default=10 * 1024 * 1024, gt=0)