from __future__ import annotations import asyncio import socket import uuid from datetime import UTC, datetime, timedelta 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 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: task = await self.repository.claim(self.owner) if not task: return False 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"]) return True 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"] ) rule = detect_format(body, attachment.mime_type) if not rule: malware = await self.antivirus.scan(one_chunk(body)) rule = "file.malware_detected" if malware else None 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: 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: await self.repository.finish( task.id, self.owner, task.lease_generation, allow=False, rule_id="file.object_changed", ) except DependencyFailure: await self.repository.retry_or_fail(task, row.config["task"]["max_attempts"]) finally: stop.set() await heartbeat return True 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() 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, ) try: async with asyncio.TaskGroup() as group: for _ in range(settings.worker_concurrency): group.create_task(worker.loop()) finally: await engine.dispose() def run() -> None: asyncio.run(serve()) if __name__ == "__main__": run()