refactor(evidence): remove legacy Python layout (#31)
This commit is contained in:
@@ -0,0 +1,7 @@
|
||||
"""Evidence-owned source adapter implementations."""
|
||||
|
||||
from tht.evidence.adapters.filesystem import FilesystemEvidenceSource
|
||||
from tht.evidence.adapters.http import HttpManifestEvidenceSource
|
||||
from tht.evidence.adapters.s3 import S3EvidenceSource
|
||||
|
||||
__all__ = ["FilesystemEvidenceSource", "HttpManifestEvidenceSource", "S3EvidenceSource"]
|
||||
@@ -0,0 +1,148 @@
|
||||
"""Evidence-owned, race-safe filesystem source."""
|
||||
|
||||
import hashlib
|
||||
import os
|
||||
import stat
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path, PurePosixPath
|
||||
from urllib.parse import unquote, urlsplit
|
||||
|
||||
from tht.evidence.contracts import (
|
||||
AcquiredDocument,
|
||||
EvidenceSourceError,
|
||||
EvidenceSourceErrorCategory,
|
||||
SourceObject,
|
||||
)
|
||||
|
||||
|
||||
class FilesystemEvidenceSource:
|
||||
def __init__(
|
||||
self,
|
||||
root: Path | str,
|
||||
*,
|
||||
patterns: tuple[str, ...] | list[str] = ("**/*.md",),
|
||||
max_bytes: int = 10 * 1024 * 1024,
|
||||
) -> None:
|
||||
if max_bytes < 1:
|
||||
raise ValueError("max_bytes must be positive")
|
||||
if not patterns or any(
|
||||
not pattern or Path(pattern).is_absolute() or ".." in Path(pattern).parts
|
||||
for pattern in patterns
|
||||
):
|
||||
raise ValueError("at least one non-empty discovery pattern is required")
|
||||
try:
|
||||
self.root = Path(root).expanduser().resolve(strict=True)
|
||||
self._root_fd = os.open(
|
||||
self.root,
|
||||
os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW | os.O_CLOEXEC,
|
||||
)
|
||||
except OSError as error:
|
||||
raise ValueError("filesystem evidence root is unavailable") from error
|
||||
self.patterns = tuple(patterns)
|
||||
self.max_bytes = max_bytes
|
||||
|
||||
def __del__(self):
|
||||
root_fd = getattr(self, "_root_fd", None)
|
||||
if root_fd is not None:
|
||||
try:
|
||||
os.close(root_fd)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
def _safe_error(operation: str, *, transient: bool = False, **details):
|
||||
return EvidenceSourceError(
|
||||
"filesystem source operation failed",
|
||||
category=(
|
||||
EvidenceSourceErrorCategory.TRANSIENT
|
||||
if transient
|
||||
else EvidenceSourceErrorCategory.PERMANENT
|
||||
),
|
||||
details={"operation": operation, **details},
|
||||
)
|
||||
|
||||
def _open_read(self, relative: PurePosixPath) -> tuple[bytes, os.stat_result]:
|
||||
parts = relative.parts
|
||||
if not parts or any(part in {"", ".", ".."} for part in parts):
|
||||
raise self._safe_error("path_validation")
|
||||
directory_fd = os.dup(self._root_fd)
|
||||
file_fd = None
|
||||
try:
|
||||
for component in parts[:-1]:
|
||||
next_fd = os.open(
|
||||
component,
|
||||
os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW | os.O_CLOEXEC,
|
||||
dir_fd=directory_fd,
|
||||
)
|
||||
os.close(directory_fd)
|
||||
directory_fd = next_fd
|
||||
file_fd = os.open(
|
||||
parts[-1],
|
||||
os.O_RDONLY | os.O_NOFOLLOW | os.O_CLOEXEC,
|
||||
dir_fd=directory_fd,
|
||||
)
|
||||
file_stat = os.fstat(file_fd)
|
||||
if not stat.S_ISREG(file_stat.st_mode):
|
||||
raise self._safe_error("path_validation")
|
||||
if file_stat.st_size > self.max_bytes:
|
||||
raise self._safe_error("read", limit_bytes=self.max_bytes)
|
||||
content = bytearray()
|
||||
while len(content) <= self.max_bytes:
|
||||
chunk = os.read(file_fd, min(64 * 1024, self.max_bytes + 1 - len(content)))
|
||||
if not chunk:
|
||||
break
|
||||
content.extend(chunk)
|
||||
if len(content) > self.max_bytes:
|
||||
raise self._safe_error("read", limit_bytes=self.max_bytes)
|
||||
return bytes(content), file_stat
|
||||
except EvidenceSourceError:
|
||||
raise
|
||||
except OSError as error:
|
||||
raise self._safe_error("open") from error
|
||||
finally:
|
||||
if file_fd is not None:
|
||||
os.close(file_fd)
|
||||
os.close(directory_fd)
|
||||
|
||||
def _item(
|
||||
self, relative: PurePosixPath, content: bytes, file_stat: os.stat_result
|
||||
) -> SourceObject:
|
||||
relative_text = relative.as_posix()
|
||||
return SourceObject(
|
||||
source_id=f"filesystem:{hashlib.sha256(relative_text.encode()).hexdigest()}",
|
||||
uri=(self.root / relative_text).as_uri(),
|
||||
fingerprint=f"sha256:{hashlib.sha256(content).hexdigest()}",
|
||||
modified_at=datetime.fromtimestamp(file_stat.st_mtime, tz=UTC),
|
||||
metadata={"relative_path": relative_text},
|
||||
)
|
||||
|
||||
def discover(self):
|
||||
candidates = {
|
||||
path.relative_to(self.root).as_posix()
|
||||
for pattern in self.patterns
|
||||
for path in self.root.glob(pattern)
|
||||
}
|
||||
for relative_text in sorted(candidates):
|
||||
relative = PurePosixPath(relative_text)
|
||||
content, file_stat = self._open_read(relative)
|
||||
yield self._item(relative, content, file_stat)
|
||||
|
||||
def acquire(self, item: SourceObject) -> AcquiredDocument:
|
||||
parsed = urlsplit(item.uri)
|
||||
if parsed.scheme != "file" or parsed.netloc or parsed.query or parsed.fragment:
|
||||
raise self._safe_error("acquire")
|
||||
try:
|
||||
relative = Path(unquote(parsed.path)).relative_to(self.root)
|
||||
except ValueError as error:
|
||||
raise self._safe_error("acquire") from error
|
||||
pure_relative = PurePosixPath(relative.as_posix())
|
||||
content, file_stat = self._open_read(pure_relative)
|
||||
expected = self._item(pure_relative, content, file_stat)
|
||||
if item.source_id != expected.source_id or item.fingerprint != expected.fingerprint:
|
||||
raise self._safe_error("acquire")
|
||||
return AcquiredDocument(
|
||||
source=expected,
|
||||
content=content,
|
||||
media_type="text/markdown" if relative.suffix.lower() == ".md" else None,
|
||||
acquired_at=datetime.now(UTC),
|
||||
)
|
||||
@@ -0,0 +1,287 @@
|
||||
"""Evidence-owned explicit-manifest HTTP source with SSRF-safe bounded acquisition."""
|
||||
|
||||
import hashlib
|
||||
import ipaddress
|
||||
import socket
|
||||
from collections import OrderedDict
|
||||
from datetime import UTC, datetime
|
||||
from email.utils import parsedate_to_datetime
|
||||
from urllib.parse import urljoin, urlsplit
|
||||
|
||||
import requests
|
||||
|
||||
from tht.evidence.contracts import (
|
||||
AcquiredDocument,
|
||||
EvidenceSourceError,
|
||||
EvidenceSourceErrorCategory,
|
||||
SourceObject,
|
||||
canonical_provenance_uri,
|
||||
)
|
||||
|
||||
|
||||
class HttpManifestEvidenceSource:
|
||||
def __init__(
|
||||
self,
|
||||
urls: list[str] | tuple[str, ...],
|
||||
*,
|
||||
connect_timeout: float = 5,
|
||||
read_timeout: float = 30,
|
||||
max_bytes: int = 10 * 1024 * 1024,
|
||||
max_redirects: int = 5,
|
||||
allow_private_hosts: bool = False,
|
||||
max_cache_bytes: int = 64 * 1024 * 1024,
|
||||
) -> None:
|
||||
if not urls:
|
||||
raise ValueError("HTTP evidence manifest must contain at least one URL")
|
||||
if (
|
||||
connect_timeout <= 0
|
||||
or read_timeout <= 0
|
||||
or max_bytes < 1
|
||||
or max_redirects < 0
|
||||
or max_cache_bytes < 1
|
||||
):
|
||||
raise ValueError("HTTP evidence limits must be positive")
|
||||
self._transport_by_uri: dict[str, str] = {}
|
||||
for url in urls:
|
||||
self._validate_url_shape(url)
|
||||
provenance = canonical_provenance_uri(url)
|
||||
if provenance in self._transport_by_uri:
|
||||
raise ValueError("HTTP evidence manifest contains duplicate canonical provenance")
|
||||
self._transport_by_uri[provenance] = url
|
||||
self.connect_timeout = connect_timeout
|
||||
self.read_timeout = read_timeout
|
||||
self.max_bytes = max_bytes
|
||||
self.max_redirects = max_redirects
|
||||
self.allow_private_hosts = allow_private_hosts
|
||||
self.max_cache_bytes = max_cache_bytes
|
||||
self._session = requests.Session()
|
||||
self._session.trust_env = False
|
||||
self._cache: OrderedDict[str, AcquiredDocument] = OrderedDict()
|
||||
# provenance -> (exact final effective URL, ETag, Last-Modified)
|
||||
self._validators: dict[str, tuple[str, str | None, str | None]] = {}
|
||||
self._cache_bytes = 0
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"HttpManifestEvidenceSource(objects={len(self._transport_by_uri)})"
|
||||
|
||||
@staticmethod
|
||||
def _safe_error(operation: str, *, transient: bool = False, **details):
|
||||
return EvidenceSourceError(
|
||||
"HTTP source operation failed",
|
||||
category=(
|
||||
EvidenceSourceErrorCategory.TRANSIENT
|
||||
if transient
|
||||
else EvidenceSourceErrorCategory.PERMANENT
|
||||
),
|
||||
details={"operation": operation, **details},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _validate_url_shape(url: str) -> None:
|
||||
parsed = urlsplit(url)
|
||||
if parsed.scheme not in {"http", "https"} or not parsed.hostname:
|
||||
raise ValueError("HTTP evidence URLs must use http or https")
|
||||
if parsed.username is not None or parsed.password is not None:
|
||||
raise ValueError("HTTP evidence URLs must not contain userinfo credentials")
|
||||
|
||||
@staticmethod
|
||||
def _source_id(uri: str) -> str:
|
||||
return f"http:{hashlib.sha256(uri.encode()).hexdigest()}"
|
||||
|
||||
@staticmethod
|
||||
def _normalized_ip(value: str) -> ipaddress.IPv4Address | ipaddress.IPv6Address:
|
||||
address = ipaddress.ip_address(value.split("%", 1)[0])
|
||||
if isinstance(address, ipaddress.IPv6Address) and address.ipv4_mapped:
|
||||
return address.ipv4_mapped
|
||||
return address
|
||||
|
||||
def _resolve_allowed(self, url: str) -> set[ipaddress.IPv4Address | ipaddress.IPv6Address]:
|
||||
try:
|
||||
self._validate_url_shape(url)
|
||||
except ValueError as error:
|
||||
raise self._safe_error("url_validation") from error
|
||||
if self.allow_private_hosts:
|
||||
return set()
|
||||
parsed = urlsplit(url)
|
||||
port = parsed.port or (443 if parsed.scheme == "https" else 80)
|
||||
try:
|
||||
rows = socket.getaddrinfo(parsed.hostname, port, type=socket.SOCK_STREAM)
|
||||
addresses = {self._normalized_ip(row[4][0]) for row in rows}
|
||||
except (OSError, ValueError) as error:
|
||||
raise self._safe_error("resolution", transient=True) from error
|
||||
if not addresses:
|
||||
raise self._safe_error("resolution", transient=True)
|
||||
# Reject the entire answer set if any address is private/reserved. Choosing only a public
|
||||
# member would leave DNS ordering as a policy bypass.
|
||||
if any(not address.is_global for address in addresses):
|
||||
raise self._safe_error("network_policy")
|
||||
return addresses
|
||||
|
||||
def _verify_peer(
|
||||
self,
|
||||
response,
|
||||
allowed: set[ipaddress.IPv4Address | ipaddress.IPv6Address],
|
||||
) -> None:
|
||||
if self.allow_private_hosts:
|
||||
return
|
||||
try:
|
||||
connection = response.raw._connection
|
||||
peer = self._normalized_ip(connection.sock.getpeername()[0])
|
||||
except (AttributeError, OSError, TypeError, ValueError) as error:
|
||||
raise self._safe_error("peer_validation", transient=True) from error
|
||||
if not peer.is_global or peer not in allowed:
|
||||
raise self._safe_error("network_policy")
|
||||
|
||||
@staticmethod
|
||||
def _status_category(status: int) -> EvidenceSourceErrorCategory:
|
||||
if status in {408, 425, 429} or 500 <= status <= 599:
|
||||
return EvidenceSourceErrorCategory.TRANSIENT
|
||||
return EvidenceSourceErrorCategory.PERMANENT
|
||||
|
||||
def _conditional_headers(self, provenance: str, request_url: str) -> dict[str, str]:
|
||||
cached = self._cache.get(self._source_id(provenance))
|
||||
validators = self._validators.get(provenance)
|
||||
if cached is None or validators is None:
|
||||
return {}
|
||||
final_url, etag, last_modified = validators
|
||||
if request_url != final_url:
|
||||
return {}
|
||||
headers = {}
|
||||
if etag:
|
||||
headers["If-None-Match"] = etag
|
||||
if last_modified:
|
||||
headers["If-Modified-Since"] = last_modified
|
||||
return headers
|
||||
|
||||
def _remember(
|
||||
self,
|
||||
provenance: str,
|
||||
final_url: str,
|
||||
document: AcquiredDocument,
|
||||
validators: tuple[str | None, str | None],
|
||||
) -> None:
|
||||
source_id = document.source.source_id
|
||||
old = self._cache.pop(source_id, None)
|
||||
if old is not None:
|
||||
self._cache_bytes -= len(old.content)
|
||||
self._cache[source_id] = document
|
||||
self._cache_bytes += len(document.content)
|
||||
self._validators[provenance] = (final_url, *validators)
|
||||
while self._cache and self._cache_bytes > self.max_cache_bytes:
|
||||
evicted_id, evicted = self._cache.popitem(last=False)
|
||||
self._cache_bytes -= len(evicted.content)
|
||||
for uri in tuple(self._validators):
|
||||
if self._source_id(uri) == evicted_id:
|
||||
del self._validators[uri]
|
||||
|
||||
def _download(self, transport_url: str, provenance: str) -> AcquiredDocument:
|
||||
current = transport_url
|
||||
try:
|
||||
for redirect_count in range(self.max_redirects + 1):
|
||||
headers = self._conditional_headers(provenance, current)
|
||||
allowed = self._resolve_allowed(current)
|
||||
response = None
|
||||
try:
|
||||
response = self._session.get(
|
||||
current,
|
||||
headers=headers,
|
||||
stream=True,
|
||||
allow_redirects=False,
|
||||
timeout=(self.connect_timeout, self.read_timeout),
|
||||
)
|
||||
self._verify_peer(response, allowed)
|
||||
if response.is_redirect:
|
||||
if redirect_count == self.max_redirects:
|
||||
raise self._safe_error("redirect")
|
||||
destination = urljoin(current, response.headers.get("Location", ""))
|
||||
try:
|
||||
self._validate_url_shape(destination)
|
||||
except ValueError as error:
|
||||
raise self._safe_error("redirect") from error
|
||||
# The next iteration binds validators to the exact destination URL.
|
||||
current = destination
|
||||
continue
|
||||
if response.status_code == 304:
|
||||
cached = self._cache.get(self._source_id(provenance))
|
||||
binding = self._validators.get(provenance)
|
||||
if (
|
||||
cached is None
|
||||
or not headers
|
||||
or binding is None
|
||||
or binding[0] != current
|
||||
):
|
||||
raise self._safe_error("conditional_response")
|
||||
self._cache.move_to_end(cached.source.source_id)
|
||||
return cached
|
||||
if not 200 <= response.status_code <= 299:
|
||||
raise EvidenceSourceError(
|
||||
"HTTP status failure",
|
||||
category=self._status_category(response.status_code),
|
||||
details={"operation": "download", "status": response.status_code},
|
||||
)
|
||||
length = response.headers.get("Content-Length")
|
||||
if length is not None and int(length) > self.max_bytes:
|
||||
raise self._safe_error("download", limit_bytes=self.max_bytes)
|
||||
content = bytearray()
|
||||
for chunk in response.iter_content(
|
||||
chunk_size=min(64 * 1024, self.max_bytes + 1)
|
||||
):
|
||||
content.extend(chunk)
|
||||
if len(content) > self.max_bytes:
|
||||
raise self._safe_error("download", limit_bytes=self.max_bytes)
|
||||
etag = response.headers.get("ETag")
|
||||
last_modified = response.headers.get("Last-Modified")
|
||||
media_type = (
|
||||
response.headers.get("Content-Type", "").split(";", 1)[0] or None
|
||||
)
|
||||
break
|
||||
finally:
|
||||
if response is not None:
|
||||
response.close()
|
||||
except EvidenceSourceError:
|
||||
raise
|
||||
except (requests.Timeout, requests.ConnectionError, TimeoutError) as error:
|
||||
raise self._safe_error("download", transient=True) from error
|
||||
except requests.RequestException as error:
|
||||
raise self._safe_error("download", transient=True) from error
|
||||
except (OSError, ValueError) as error:
|
||||
raise self._safe_error("download") from error
|
||||
|
||||
modified_at = None
|
||||
if etag:
|
||||
fingerprint = f"etag:{hashlib.sha256(etag.encode()).hexdigest()}"
|
||||
elif last_modified:
|
||||
try:
|
||||
modified_at = parsedate_to_datetime(last_modified).astimezone(UTC)
|
||||
fingerprint = f"last-modified:{int(modified_at.timestamp())}"
|
||||
except (TypeError, ValueError, OverflowError):
|
||||
fingerprint = f"sha256:{hashlib.sha256(content).hexdigest()}"
|
||||
else:
|
||||
fingerprint = f"sha256:{hashlib.sha256(content).hexdigest()}"
|
||||
item = SourceObject(
|
||||
source_id=self._source_id(provenance),
|
||||
uri=provenance,
|
||||
fingerprint=fingerprint,
|
||||
modified_at=modified_at,
|
||||
)
|
||||
document = AcquiredDocument(
|
||||
source=item,
|
||||
content=bytes(content),
|
||||
media_type=media_type,
|
||||
acquired_at=datetime.now(UTC),
|
||||
)
|
||||
self._remember(provenance, current, document, (etag, last_modified))
|
||||
return document
|
||||
|
||||
def discover(self):
|
||||
for provenance in sorted(self._transport_by_uri):
|
||||
yield self._download(self._transport_by_uri[provenance], provenance).source
|
||||
|
||||
def acquire(self, item: SourceObject) -> AcquiredDocument:
|
||||
transport = self._transport_by_uri.get(item.uri)
|
||||
if transport is None or item.source_id != self._source_id(item.uri):
|
||||
raise self._safe_error("acquire")
|
||||
document = self._download(transport, item.uri)
|
||||
if document.source.fingerprint != item.fingerprint:
|
||||
raise self._safe_error("acquire")
|
||||
return document
|
||||
@@ -0,0 +1,152 @@
|
||||
"""Evidence-owned bounded S3-compatible source using the supported boto3 client."""
|
||||
|
||||
import hashlib
|
||||
import ipaddress
|
||||
import re
|
||||
from datetime import UTC, datetime
|
||||
from urllib.parse import quote, urlsplit
|
||||
|
||||
from tht.evidence.contracts 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:
|
||||
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("/") or len(prefix.encode()) > 1024
|
||||
or any(ord(char) < 32 or ord(char) == 127 for char in prefix)):
|
||||
raise ValueError("S3 prefix is invalid")
|
||||
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")
|
||||
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
|
||||
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[SourceObject, str]] = {}
|
||||
|
||||
@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:
|
||||
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
|
||||
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 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()
|
||||
modified = row.get("LastModified")
|
||||
if modified is not None:
|
||||
modified = modified.astimezone(UTC)
|
||||
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:
|
||||
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 or item != binding[0]:
|
||||
raise self._error("acquire")
|
||||
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}
|
||||
body = None
|
||||
try:
|
||||
response = self._client.get_object(**kwargs)
|
||||
body = response["Body"]
|
||||
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)
|
||||
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()
|
||||
Reference in New Issue
Block a user