Files

437 lines
16 KiB
Python

from __future__ import annotations
import asyncio
import uuid
from datetime import UTC, datetime, timedelta
from time import perf_counter
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 Antivirus, DependencyFailure, 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.telemetry import (
current_trace_fields,
record_cache,
record_check,
record_dependency,
record_file_rejection,
record_runtime_state,
)
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",
antivirus: Antivirus | None = None,
) -> 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.antivirus = antivirus
self._antivirus_checked_at = 0.0
self._antivirus_check_lock = asyncio.Lock()
self.rate_limiter = ConservativeRateLimiter(
config.document["rate"]["text_rps"], config.document["rate"]["file_rps"]
)
record_runtime_state("mock" if mode.mock else "standard", config.version)
async def refresh_antivirus(self) -> bool:
if self.mode.mock:
self.files_ready = False
return False
if self.antivirus is None:
return self.files_ready
if perf_counter() - self._antivirus_checked_at < 5:
return self.files_ready
async with self._antivirus_check_lock:
if perf_counter() - self._antivirus_checked_at < 5:
return self.files_ready
try:
status = await self.antivirus.status()
maximum_age = timedelta(
hours=self.config.document["antivirus"]["max_signature_age_hours"]
)
if datetime.now(UTC) - status.databases_date > maximum_age:
raise DependencyFailure("KESL databases are stale")
self.signatures_version = status.signatures_version
self.files_ready = True
except DependencyFailure:
self.signatures_version = "unavailable"
self.files_ready = False
self._antivirus_checked_at = perf_counter()
return self.files_ready
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:
started = perf_counter()
outcome = "error"
try:
result = await self._check(request)
outcome = "pending" if isinstance(result, Pending) else result.verdict
return result
finally:
record_check(
request.content_kind,
"mock" if self.mode.mock else "standard",
outcome,
(perf_counter() - started) * 1000,
self.config.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
)
record_cache("text_rules", cache is not None)
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
)
record_cache("link_verdict", cached_link is not None)
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"]
)
record_dependency("dns", "resolve", "success")
except DnsError as exc:
record_dependency("dns", "resolve", "error")
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 await self.refresh_antivirus():
raise CapabilityUnavailable("files")
rule = validate_metadata(
request.attachment,
self.config.detector,
set(self.config.document["file_policy"]["enabled_mime_types"]),
)
if rule:
record_file_rejection(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
)
record_cache("file_verdict", cached is not None)
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()
trace_id, span_id, trace_flags, tracestate = current_trace_fields()
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="kesl",
signatures_version=self.signatures_version,
origin_trace_id=trace_id,
origin_span_id=span_id,
origin_trace_flags=trace_flags,
origin_tracestate=tracestate,
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:
record_dependency("postgresql", "queue_capacity", "full")
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"]),
)
)