368 lines
14 KiB
Python
368 lines
14 KiB
Python
from __future__ import annotations
|
|
|
|
import uuid
|
|
from datetime import UTC, datetime, timedelta
|
|
|
|
from app.config import ActiveConfig
|
|
from app.contracts import CheckRequest, FileCheck, Pending, TextCheck, Verdict
|
|
from app.db import (
|
|
LinkVerdictCache,
|
|
SafetyAudit,
|
|
SafetyRequest,
|
|
SafetyTask,
|
|
TaskStatus,
|
|
TextRulesCache,
|
|
)
|
|
from app.file_pipeline import validate_metadata
|
|
from app.fingerprint import fingerprint
|
|
from app.normalization import normalize_text
|
|
from app.rate_limit import ConservativeRateLimiter
|
|
from app.repository import QueueFull, Repository
|
|
from app.settings import EmergencyMode
|
|
from app.url_policy import DnsError, Resolver, canonicalize, check_url, extract_urls
|
|
|
|
|
|
class CapabilityUnavailable(RuntimeError):
|
|
def __init__(self, category: str) -> None:
|
|
self.category = category
|
|
|
|
|
|
class TaskFailed(RuntimeError):
|
|
def __init__(self, task_id: uuid.UUID) -> None:
|
|
self.task_id = task_id
|
|
|
|
|
|
class SafetyService:
|
|
def __init__(
|
|
self,
|
|
repository: Repository,
|
|
config: ActiveConfig,
|
|
mode: EmergencyMode,
|
|
resolver: Resolver,
|
|
*,
|
|
links_ready: bool = True,
|
|
files_ready: bool = True,
|
|
signatures_version: str = "unverified",
|
|
) -> None:
|
|
self.repository = repository
|
|
self.config = config
|
|
self.mode = mode
|
|
self.resolver = resolver
|
|
self.links_ready = links_ready
|
|
self.files_ready = files_ready
|
|
self.signatures_version = signatures_version
|
|
self.rate_limiter = ConservativeRateLimiter(
|
|
config.document["rate"]["text_rps"], config.document["rate"]["file_rps"]
|
|
)
|
|
|
|
def _verdict(
|
|
self,
|
|
allow: bool,
|
|
mode: str,
|
|
rule: str,
|
|
rules_version: str,
|
|
*,
|
|
config_version: int | None = None,
|
|
) -> Verdict:
|
|
return Verdict(
|
|
verdict="allow" if allow else "deny",
|
|
processing_mode=mode,
|
|
config_version=self.config.version if config_version is None else config_version,
|
|
rule_id=rule,
|
|
reason_code=None if allow else "message_blocked",
|
|
rules_version=rules_version,
|
|
)
|
|
|
|
async def check(self, request: CheckRequest) -> Verdict | Pending:
|
|
digest = fingerprint(request)
|
|
existing = await self.repository.get_request(request.message_id)
|
|
if existing:
|
|
if existing.request_fingerprint != digest:
|
|
from app.repository import ConflictError
|
|
|
|
raise ConflictError
|
|
return await self._replay(existing)
|
|
await self.rate_limiter.acquire(request.content_kind)
|
|
if self.mode.mock:
|
|
free = self.mode.text_free if request.content_kind == "text" else self.mode.file_free
|
|
verdict = self._verdict(
|
|
free,
|
|
"mock",
|
|
"safety.mock_forced_allow" if free else "safety.mock_forced_deny",
|
|
"mock",
|
|
)
|
|
return await self._persist_sync(request, digest, verdict)
|
|
if isinstance(request, TextCheck):
|
|
return await self._check_text(request, digest)
|
|
return await self._check_file(request, digest)
|
|
|
|
async def _check_text(self, request: TextCheck, digest: bytes) -> Verdict:
|
|
normalized = normalize_text(request.text)
|
|
cache = await self.repository.text_cache(
|
|
normalized.analysis_sha256, self.config.rules_version
|
|
)
|
|
if cache:
|
|
deny_rule = cache.deny_rule_id
|
|
else:
|
|
result = self.config.rules.evaluate(normalized.analysis)
|
|
deny_rule = result.deny_rule
|
|
now = datetime.now(UTC)
|
|
await self.repository.put_text_cache(
|
|
TextRulesCache(
|
|
analysis_sha256=normalized.analysis_sha256,
|
|
rules_version=self.config.rules_version,
|
|
result="deny" if deny_rule else "allow",
|
|
deny_rule_id=deny_rule,
|
|
monitor_rule_ids=list(result.monitor_rules),
|
|
normalization_flags=list(normalized.flags),
|
|
created_at=now,
|
|
expires_at=now
|
|
+ timedelta(seconds=self.config.document["cache"]["text_rule_ttl_sec"]),
|
|
)
|
|
)
|
|
if deny_rule:
|
|
return await self._persist_sync(
|
|
request,
|
|
digest,
|
|
self._verdict(False, "standard", deny_rule, self.config.rules_version),
|
|
)
|
|
urls = extract_urls(
|
|
normalized.analysis,
|
|
maximum=self.config.document["link"]["max_per_message"],
|
|
max_length=self.config.document["link"]["url_max_length"],
|
|
)
|
|
if urls and not self.links_ready:
|
|
raise CapabilityUnavailable("dns")
|
|
for raw in urls:
|
|
try:
|
|
canonical = canonicalize(raw)
|
|
except PermissionError as exc:
|
|
return await self._persist_sync(
|
|
request,
|
|
digest,
|
|
self._verdict(False, "standard", str(exc), self.config.rules_version),
|
|
)
|
|
cached_link = await self.repository.link_cache(
|
|
canonical.digest, self.config.rules_version, self.config.version
|
|
)
|
|
if cached_link and cached_link.verdict == "deny":
|
|
return await self._persist_sync(
|
|
request,
|
|
digest,
|
|
self._verdict(
|
|
False,
|
|
"standard",
|
|
cached_link.rule_id or "url.malformed",
|
|
self.config.rules_version,
|
|
),
|
|
)
|
|
try:
|
|
_, rule = await check_url(
|
|
raw, self.resolver, self.config.document["link"]["dns_lookup_timeout_sec"]
|
|
)
|
|
except DnsError as exc:
|
|
raise CapabilityUnavailable("dns") from exc
|
|
if rule != "url.nxdomain" and not cached_link:
|
|
now = datetime.now(UTC)
|
|
await self.repository.put_link_cache(
|
|
LinkVerdictCache(
|
|
canonical_url_sha256=canonical.digest,
|
|
rules_version=self.config.rules_version,
|
|
config_version=self.config.version,
|
|
verdict="deny" if rule else "allow",
|
|
rule_id=rule,
|
|
reason_code="message_blocked" if rule else None,
|
|
first_seen_at=now,
|
|
last_seen_at=now,
|
|
hit_count=1,
|
|
expires_at=now
|
|
+ timedelta(seconds=self.config.document["cache"]["link_ttl_sec"]),
|
|
)
|
|
)
|
|
if rule and rule != "url.nxdomain":
|
|
return await self._persist_sync(
|
|
request,
|
|
digest,
|
|
self._verdict(False, "standard", rule, self.config.rules_version),
|
|
)
|
|
return await self._persist_sync(
|
|
request,
|
|
digest,
|
|
self._verdict(True, "standard", "safety.all_checks_passed", self.config.rules_version),
|
|
)
|
|
|
|
async def _check_file(self, request: FileCheck, digest: bytes) -> Verdict | Pending:
|
|
if not self.files_ready:
|
|
raise CapabilityUnavailable("files")
|
|
rule = validate_metadata(
|
|
request.attachment,
|
|
self.config.detector,
|
|
set(self.config.document["file_policy"]["enabled_mime_types"]),
|
|
)
|
|
if rule:
|
|
return await self._persist_sync(
|
|
request, digest, self._verdict(False, "standard", rule, self.config.rules_version)
|
|
)
|
|
content_digest = bytes.fromhex(request.attachment.checksum[7:])
|
|
cached = await self.repository.file_cache(
|
|
content_digest, self.config, self.signatures_version
|
|
)
|
|
if cached:
|
|
return await self._persist_sync(
|
|
request,
|
|
digest,
|
|
self._verdict(
|
|
cached.verdict == "allow",
|
|
"standard",
|
|
cached.rule_id,
|
|
cached.rules_version,
|
|
),
|
|
)
|
|
now = datetime.now(UTC)
|
|
task_id = uuid.uuid4()
|
|
task = SafetyTask(
|
|
id=task_id,
|
|
message_id=request.message_id,
|
|
attachment_id=request.attachment.attachment_id,
|
|
request_fingerprint=digest,
|
|
content_sha256=content_digest,
|
|
processing_mode="standard",
|
|
config_version=self.config.version,
|
|
status=TaskStatus.pending,
|
|
attempt_count=0,
|
|
lease_generation=0,
|
|
expires_at=now
|
|
+ timedelta(seconds=self.config.document["task"]["execution_deadline_sec"]),
|
|
quarantine_object_key=request.attachment.quarantine_object_key,
|
|
quarantine_version_id=request.attachment.quarantine_version_id,
|
|
quarantine_etag=request.attachment.quarantine_etag,
|
|
declared_mime=request.attachment.mime_type,
|
|
declared_size_bytes=request.attachment.size_bytes,
|
|
declared_checksum=request.attachment.checksum,
|
|
rules_version=self.config.rules_version,
|
|
detector_version=self.config.detector.version,
|
|
scanner_engine="clamav",
|
|
signatures_version=self.signatures_version,
|
|
created_at=now,
|
|
updated_at=now,
|
|
)
|
|
row = SafetyRequest(
|
|
message_id=request.message_id,
|
|
request_fingerprint=digest,
|
|
processing_mode="standard",
|
|
config_version=self.config.version,
|
|
verdict="pending",
|
|
task_id=task_id,
|
|
rules_version=self.config.rules_version,
|
|
created_at=now,
|
|
purge_after=now + timedelta(days=30),
|
|
)
|
|
try:
|
|
stored, created = await self.repository.reserve_request(
|
|
row, task, max_pending=self.config.document["task"]["max_pending"]
|
|
)
|
|
except QueueFull as exc:
|
|
raise CapabilityUnavailable("queue_capacity") from exc
|
|
if created:
|
|
await self._audit(
|
|
request.message_id,
|
|
"task_created",
|
|
"standard",
|
|
"pending",
|
|
None,
|
|
task_id=task.id,
|
|
)
|
|
return await self._replay(stored)
|
|
|
|
async def _persist_sync(
|
|
self, request: CheckRequest, digest: bytes, verdict: Verdict
|
|
) -> Verdict:
|
|
now = datetime.now(UTC)
|
|
row = SafetyRequest(
|
|
message_id=request.message_id,
|
|
request_fingerprint=digest,
|
|
processing_mode=verdict.processing_mode,
|
|
config_version=verdict.config_version,
|
|
verdict=verdict.verdict,
|
|
rule_id=verdict.rule_id,
|
|
reason_code=verdict.reason_code,
|
|
rules_version=verdict.rules_version,
|
|
created_at=now,
|
|
purge_after=now + timedelta(days=30),
|
|
)
|
|
stored, created = await self.repository.reserve_request(row)
|
|
if created:
|
|
event = (
|
|
f"mock_forced_{verdict.verdict}"
|
|
if verdict.processing_mode == "mock"
|
|
else ("rule_hit" if verdict.verdict == "deny" else "received")
|
|
)
|
|
await self._audit(
|
|
request.message_id,
|
|
event,
|
|
verdict.processing_mode,
|
|
verdict.verdict,
|
|
verdict.rule_id,
|
|
)
|
|
replay = await self._replay(stored)
|
|
assert isinstance(replay, Verdict)
|
|
return replay
|
|
|
|
async def _replay(self, row: SafetyRequest) -> Verdict | Pending:
|
|
if row.verdict == "pending":
|
|
assert row.task_id
|
|
task = await self.repository.task(row.task_id)
|
|
if task and task.status in {TaskStatus.allowed, TaskStatus.denied}:
|
|
return self._verdict(
|
|
task.status == TaskStatus.allowed,
|
|
task.processing_mode,
|
|
task.rule_id or "safety.all_checks_passed",
|
|
task.rules_version,
|
|
config_version=task.config_version,
|
|
)
|
|
if task and task.status == TaskStatus.failed:
|
|
raise TaskFailed(task.id)
|
|
assert task
|
|
return Pending(
|
|
config_version=task.config_version,
|
|
task_id=task.id,
|
|
expires_at=task.expires_at,
|
|
rules_version=task.rules_version,
|
|
)
|
|
return Verdict(
|
|
verdict=row.verdict,
|
|
processing_mode=row.processing_mode,
|
|
config_version=row.config_version,
|
|
rule_id=row.rule_id or "safety.all_checks_passed",
|
|
reason_code=row.reason_code,
|
|
rules_version=row.rules_version,
|
|
)
|
|
|
|
async def _audit(
|
|
self,
|
|
message_id: uuid.UUID,
|
|
event: str,
|
|
mode: str,
|
|
verdict: str,
|
|
rule_id: str | None,
|
|
*,
|
|
task_id: uuid.UUID | None = None,
|
|
) -> None:
|
|
now = datetime.now(UTC)
|
|
await self.repository.audit(
|
|
SafetyAudit(
|
|
id=uuid.uuid4(),
|
|
message_id=message_id,
|
|
task_id=task_id,
|
|
event=event,
|
|
processing_mode=mode,
|
|
config_version=self.config.version,
|
|
verdict=verdict,
|
|
rule_id=rule_id,
|
|
rules_version=self.config.rules_version if mode == "standard" else "mock",
|
|
normalization_flags=[],
|
|
created_at=now,
|
|
purge_after=now + timedelta(days=self.config.document["retention"]["audit_days"]),
|
|
)
|
|
)
|