Проект разделен на два репозитория
This commit is contained in:
@@ -0,0 +1,173 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user