126 lines
6.4 KiB
Python
126 lines
6.4 KiB
Python
"""Bounded S3-compatible Evidence source using the supported boto3 client."""
|
|
|
|
import hashlib
|
|
import ipaddress
|
|
import socket
|
|
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,
|
|
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)):
|
|
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:
|
|
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")
|
|
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
|
|
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)
|
|
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, version, etag = row["Key"], row.get("VersionId"), row.get("ETag")
|
|
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}"
|
|
source_id = "s3:" + hashlib.sha256(uri.encode()).hexdigest()
|
|
self._items[source_id] = (key, version, 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()
|