Files
ThothII/harness/tht/adapters/evidence/s3.py
T

133 lines
7.0 KiB
Python

"""Bounded S3-compatible Evidence source using the supported boto3 client."""
import hashlib
import re
from datetime import UTC, datetime
from urllib.parse import quote, urlsplit
from tht.ports.evidence import (
AcquiredDocument, EvidenceSourceError, EvidenceSourceErrorCategory, SourceObject,
)
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 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 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 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,
config=BotoConfig(s3={"addressing_style": "path"}))
self._client = client
self._items: dict[str, tuple[str, str | None, str | None]] = {}
@staticmethod
def _error(operation: str, transient: bool = False):
return EvidenceSourceError("S3 source operation failed",
category=(EvidenceSourceErrorCategory.TRANSIENT if transient
else EvidenceSourceErrorCategory.PERMANENT),
details={"operation": operation})
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}):
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)):
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))})
except EvidenceSourceError:
raise
except Exception as exc:
raise self._error("list", transient=True) from exc
def acquire(self, item: SourceObject) -> AcquiredDocument:
binding = self._items.get(item.source_id)
if binding is None:
raise self._error("acquire")
key, version, _etag = binding
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 int(response.get("ContentLength", 0)) > self.max_bytes:
raise self._error("download_limit")
content = body.read(self.max_bytes + 1)
if len(content) > self.max_bytes:
raise self._error("download_limit")
return AcquiredDocument(source=item, content=content,
media_type=response.get("ContentType"),
acquired_at=datetime.now(UTC))
except EvidenceSourceError:
raise
except Exception as exc:
raise self._error("download", transient=True) from exc
finally:
if body is not None:
body.close()