from __future__ import annotations import asyncio import socket import uuid from datetime import UTC, datetime, timedelta from opentelemetry import trace from app.adapters import S3VersionReader from app.config import validate_config from app.contracts import Attachment from app.db import FileVerdictCache, SafetyAudit, engine_and_sessions from app.file_pipeline import ( ClamAvInstream, DependencyFailure, ObjectChanged, collect_and_hash, detect_format, one_chunk, ) from app.repository import Repository from app.settings import BootstrapSettings from app.telemetry import ( init_telemetry, record_dependency, record_worker, shutdown_telemetry, task_link, ) tracer = trace.get_tracer("message-safety.worker") class Worker: def __init__( self, repository: Repository, reader: S3VersionReader, antivirus: ClamAvInstream, artifacts ) -> None: self.repository, self.reader, self.antivirus, self.artifacts = ( repository, reader, antivirus, artifacts, ) self.owner = f"{socket.gethostname()}:{uuid.uuid4()}" async def once(self) -> bool: with tracer.start_as_current_span("message_safety.worker.claim"): task = await self.repository.claim(self.owner) if not task: return False link = task_link(task) with tracer.start_as_current_span( "message_safety.worker.process", links=[link] if link else (), attributes={ "messaging.operation.type": "process", "message_safety.attempt": task.attempt_count, }, ): await self._process(task) return True async def _process(self, task) -> None: task_age = (datetime.now(UTC) - task.created_at).total_seconds() row = await self.repository.config_version(task.config_version) rules, detector, digest = validate_config(row.config, self.artifacts) if digest != row.config_sha256: await self.repository.retry_or_fail(task, row.config["task"]["max_attempts"]) record_worker("config_error", task_age) return stop = asyncio.Event() heartbeat = asyncio.create_task( self._heartbeat( task.id, task.lease_generation, row.config["task"]["heartbeat_sec"], row.config["task"]["lease_sec"], stop, ) ) attachment = Attachment( attachment_id=task.attachment_id, quarantine_object_key=task.quarantine_object_key, quarantine_version_id=task.quarantine_version_id, quarantine_etag=task.quarantine_etag, mime_type=task.declared_mime, size_bytes=task.declared_size_bytes, checksum=task.declared_checksum, ) try: body, _ = await collect_and_hash( self.reader, attachment, max_size=row.config["file_policy"]["max_size_bytes"] ) record_dependency("s3", "get_object", "success") rule = detect_format(body, attachment.mime_type) if not rule: malware = await self.antivirus.scan(one_chunk(body)) record_dependency("clamav", "scan", "success") rule = "file.malware_detected" if malware else None with tracer.start_as_current_span("message_safety.worker.finalize"): finished = await self.repository.finish( task.id, self.owner, task.lease_generation, allow=rule is None, rule_id=rule or "safety.all_checks_passed", ) if finished: record_worker("allow" if rule is None else "deny", task_age) now = datetime.now(UTC) await self.repository.put_file_cache( FileVerdictCache( content_sha256=task.content_sha256, config_version=task.config_version, rules_version=task.rules_version, detector_version=task.detector_version, scanner_engine=task.scanner_engine, signatures_version=task.signatures_version, verdict="allow" if rule is None else "deny", rule_id=rule or "safety.all_checks_passed", reason_code=None if rule is None else "message_blocked", created_at=now, expires_at=now + timedelta(seconds=row.config["cache"]["file_verdict_ttl_sec"]), ) ) await self.repository.audit( SafetyAudit( id=uuid.uuid4(), message_id=task.message_id, task_id=task.id, event="scan_completed", processing_mode="standard", config_version=task.config_version, verdict="allow" if rule is None else "deny", rule_id=rule or "safety.all_checks_passed", rules_version=task.rules_version, normalization_flags=[], created_at=now, purge_after=now + timedelta(days=row.config["retention"]["audit_days"]), ) ) except ObjectChanged: record_dependency("s3", "get_object", "object_changed") with tracer.start_as_current_span("message_safety.worker.finalize"): await self.repository.finish( task.id, self.owner, task.lease_generation, allow=False, rule_id="file.object_changed", ) record_worker("deny", task_age) except DependencyFailure: record_dependency("file_pipeline", "scan", "error") await self.repository.retry_or_fail(task, row.config["task"]["max_attempts"]) record_worker("retry_or_fail", task_age) finally: stop.set() await heartbeat async def _heartbeat( self, task_id, generation: int, interval: int, lease: int, stop: asyncio.Event ) -> None: while True: try: await asyncio.wait_for(stop.wait(), interval) return except TimeoutError: if not await self.repository.heartbeat(task_id, self.owner, generation, lease): return async def loop(self) -> None: while True: if not await self.once(): await asyncio.sleep(0.5) async def serve() -> None: settings = BootstrapSettings() init_telemetry("worker") engine = None try: assert settings.database_url and settings.s3_access_key and settings.s3_secret_key engine, sessions = engine_and_sessions(settings.database_url.get_secret_value()) worker = Worker( Repository(sessions), S3VersionReader( settings.s3_endpoint_url, settings.s3_bucket, settings.s3_access_key.get_secret_value(), settings.s3_secret_key.get_secret_value(), ), ClamAvInstream(settings.clamav_host, settings.clamav_port), settings.artifacts_dir, ) async with asyncio.TaskGroup() as group: for _ in range(settings.worker_concurrency): group.create_task(worker.loop()) finally: if engine is not None: await engine.dispose() shutdown_telemetry() def run() -> None: asyncio.run(serve()) if __name__ == "__main__": run()