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"]), ) )