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