Files
han-app/VM2_services/codebase/services/message-safety/app/file_pipeline.py
T

185 lines
6.4 KiB
Python

from __future__ import annotations
import asyncio
import hashlib
import io
import json
import struct
from collections.abc import AsyncIterator
from dataclasses import dataclass
from pathlib import Path
from typing import Protocol
from opentelemetry import trace
from PIL import Image, UnidentifiedImageError
from pillow_heif import register_heif_opener
from app.contracts import Attachment
register_heif_opener()
tracer = trace.get_tracer("message-safety.dependencies")
class ObjectChanged(RuntimeError):
pass
class DependencyFailure(RuntimeError):
pass
class ObjectReader(Protocol):
async def stream(self, attachment: Attachment) -> AsyncIterator[bytes]: ...
class Antivirus(Protocol):
async def scan(self, chunks: AsyncIterator[bytes]) -> str | None: ...
@dataclass(frozen=True)
class DetectorManifest:
version: str
supported: frozenset[str]
max_size: int
@classmethod
def load(cls, path: Path) -> DetectorManifest:
raw = path.read_bytes()
data = json.loads(raw)
version = "sha256:" + hashlib.sha256(raw).hexdigest()
return cls(
version, frozenset(data["supported_mime_types"]), data["hard_limits"]["max_size_bytes"]
)
def validate_metadata(
attachment: Attachment, manifest: DetectorManifest, enabled: set[str]
) -> str | None:
if attachment.mime_type not in manifest.supported or attachment.mime_type not in enabled:
return "file.unsupported_mime"
if attachment.size_bytes > manifest.max_size:
return "file.size_limit"
return None
def detect_format(data: bytes, declared: str) -> str | None:
matches: list[str] = []
if data.startswith(b"\xff\xd8\xff") and data.endswith(b"\xff\xd9"):
matches.append("image/jpeg")
if data.startswith(b"\x89PNG\r\n\x1a\n") and b"IEND" in data[-64:]:
matches.append("image/png")
if len(data) >= 12 and data[:4] == b"RIFF" and data[8:12] == b"WEBP":
matches.append("image/webp")
if len(data) >= 12 and data[4:8] == b"ftyp":
brand = data[8:12]
if brand in {b"heic", b"heix", b"hevc", b"hevx"}:
matches.append("image/heic")
if brand in {b"mif1", b"msf1"}:
matches.append("image/heif")
if data.startswith(b"%PDF-") and b"%%EOF" in data[-1024:]:
matches.append("application/pdf")
if len(matches) != 1:
return "file.polyglot_or_ambiguous"
if matches[0] != declared:
return "file.format_mismatch"
if declared == "application/pdf":
lowered = data.lower()
if b"/encrypt" in lowered:
return "file.encrypted_content"
if any(
token in lowered
for token in (b"/javascript", b"/openaction", b"/launch", b"/xfa", b"/embeddedfile")
):
return "file.active_content"
else:
try:
with Image.open(io.BytesIO(data)) as image:
width, height = image.size
if width > 10_000 or height > 10_000 or width * height > 25_000_000:
return "file.parser_limit"
image.verify()
except (UnidentifiedImageError, OSError, ValueError):
return "file.format_mismatch"
return None
async def collect_and_hash(
reader: ObjectReader, attachment: Attachment, *, max_size: int
) -> tuple[bytes, bytes]:
digest = hashlib.sha256()
body = bytearray()
async for chunk in reader.stream(attachment):
if len(body) + len(chunk) > max_size:
raise ObjectChanged("object exceeds bounded size")
digest.update(chunk)
body.extend(chunk)
expected = bytes.fromhex(attachment.checksum.removeprefix("sha256:"))
if len(body) != attachment.size_bytes or digest.digest() != expected:
raise ObjectChanged("authoritative object metadata mismatch")
return bytes(body), digest.digest()
class ClamAvInstream:
def __init__(self, host: str, port: int, timeout: float = 45.0) -> None:
self.host, self.port, self.timeout = host, port, timeout
async def scan(self, chunks: AsyncIterator[bytes]) -> str | None:
with tracer.start_as_current_span(
"message_safety.clamav.scan",
attributes={"server.address.type": "clamav"},
):
return await self._scan(chunks)
async def _scan(self, chunks: AsyncIterator[bytes]) -> str | None:
async def operation() -> str | None:
reader, writer = await asyncio.open_connection(self.host, self.port)
try:
writer.write(b"zINSTREAM\0")
async for chunk in chunks:
writer.write(struct.pack(">I", len(chunk)) + chunk)
await writer.drain()
writer.write(struct.pack(">I", 0))
await writer.drain()
result = await reader.readuntil(b"\0")
text = result.rstrip(b"\0").decode("utf-8", "replace")
if text.endswith(" OK"):
return None
if text.endswith(" FOUND"):
return text.rsplit(": ", 1)[-1].removesuffix(" FOUND")
raise DependencyFailure("invalid ClamAV response")
finally:
writer.close()
await writer.wait_closed()
try:
return await asyncio.wait_for(operation(), self.timeout)
except (OSError, TimeoutError) as exc:
raise DependencyFailure("ClamAV unavailable") from exc
async def signatures_version(self) -> str:
with tracer.start_as_current_span("message_safety.clamav.version"):
return await self._signatures_version()
async def _signatures_version(self) -> str:
try:
reader, writer = await asyncio.wait_for(
asyncio.open_connection(self.host, self.port), 2.0
)
try:
writer.write(b"zVERSION\0")
await writer.drain()
raw = await asyncio.wait_for(reader.readuntil(b"\0"), 2.0)
finally:
writer.close()
await writer.wait_closed()
except (OSError, TimeoutError) as exc:
raise DependencyFailure("ClamAV unavailable") from exc
value = raw.rstrip(b"\0")
if not value.startswith(b"ClamAV ") or len(value) > 512:
raise DependencyFailure("invalid ClamAV version response")
return "sha256:" + hashlib.sha256(value).hexdigest()
async def one_chunk(data: bytes) -> AsyncIterator[bytes]:
yield data