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

174 lines
6.2 KiB
Python

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()