174 lines
6.2 KiB
Python
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()
|