from __future__ import annotations import uuid from datetime import UTC, datetime, timedelta from sqlalchemy import func, select, text, update from sqlalchemy.dialects.postgresql import insert from sqlalchemy.ext.asyncio import async_sessionmaker from app.db import ( ConfigVersion, FileVerdictCache, LinkVerdictCache, SafetyAudit, SafetyRequest, SafetyTask, TaskStatus, TextRulesCache, ) class ConflictError(RuntimeError): pass class QueueFull(RuntimeError): pass class Repository: def __init__(self, sessions: async_sessionmaker) -> None: self.sessions = sessions async def active_config(self) -> ConfigVersion: async with self.sessions() as session: rows = ( await session.scalars(select(ConfigVersion).where(ConfigVersion.state == "active")) ).all() if len(rows) != 1: raise RuntimeError("exactly one active config is required") return rows[0] async def config_version(self, version: int) -> ConfigVersion: async with self.sessions() as session: row = await session.scalar( select(ConfigVersion).where(ConfigVersion.version == version) ) if not row: raise RuntimeError("task config snapshot is missing") return row async def get_request(self, message_id: uuid.UUID) -> SafetyRequest | None: async with self.sessions() as session: return await session.get(SafetyRequest, message_id) async def text_cache(self, digest: bytes, rules_version: str) -> TextRulesCache | None: async with self.sessions() as session: return await session.scalar( select(TextRulesCache).where( TextRulesCache.analysis_sha256 == digest, TextRulesCache.rules_version == rules_version, TextRulesCache.expires_at > func.now(), ) ) async def put_text_cache(self, row: TextRulesCache) -> None: async with self.sessions.begin() as session: await session.execute( insert(TextRulesCache) .values( analysis_sha256=row.analysis_sha256, rules_version=row.rules_version, result=row.result, deny_rule_id=row.deny_rule_id, monitor_rule_ids=row.monitor_rule_ids, normalization_flags=row.normalization_flags, created_at=row.created_at, expires_at=row.expires_at, ) .on_conflict_do_nothing(constraint="uq_text_cache_key") ) async def file_cache( self, digest: bytes, config: object, signatures_version: str ) -> FileVerdictCache | None: async with self.sessions() as session: return await session.scalar( select(FileVerdictCache).where( FileVerdictCache.content_sha256 == digest, FileVerdictCache.config_version == config.version, FileVerdictCache.rules_version == config.rules_version, FileVerdictCache.detector_version == config.detector.version, FileVerdictCache.scanner_engine == "kesl", FileVerdictCache.signatures_version == signatures_version, FileVerdictCache.expires_at > func.now(), ) ) async def put_file_cache(self, row: FileVerdictCache) -> None: async with self.sessions.begin() as session: await session.execute( insert(FileVerdictCache) .values( content_sha256=row.content_sha256, config_version=row.config_version, rules_version=row.rules_version, detector_version=row.detector_version, scanner_engine=row.scanner_engine, signatures_version=row.signatures_version, verdict=row.verdict, rule_id=row.rule_id, reason_code=row.reason_code, created_at=row.created_at, expires_at=row.expires_at, ) .on_conflict_do_nothing(constraint="uq_file_cache_key") ) async def link_cache( self, digest: bytes, rules_version: str, config_version: int ) -> LinkVerdictCache | None: async with self.sessions.begin() as session: row = await session.scalar( select(LinkVerdictCache).where( LinkVerdictCache.canonical_url_sha256 == digest, LinkVerdictCache.rules_version == rules_version, LinkVerdictCache.config_version == config_version, LinkVerdictCache.expires_at > func.now(), ) ) if row: row.last_seen_at = datetime.now(UTC) row.hit_count += 1 return row async def put_link_cache(self, row: LinkVerdictCache) -> None: async with self.sessions.begin() as session: await session.execute( insert(LinkVerdictCache) .values( canonical_url_sha256=row.canonical_url_sha256, rules_version=row.rules_version, config_version=row.config_version, verdict=row.verdict, rule_id=row.rule_id, reason_code=row.reason_code, first_seen_at=row.first_seen_at, last_seen_at=row.last_seen_at, hit_count=row.hit_count, expires_at=row.expires_at, ) .on_conflict_do_nothing(constraint="uq_link_key") ) async def reserve_request( self, record: SafetyRequest, task: SafetyTask | None = None, *, max_pending: int | None = None, ) -> tuple[SafetyRequest, bool]: async with self.sessions.begin() as session: if task is not None: await session.execute( text( "SELECT pg_advisory_xact_lock(hashtext('message_safety.pending_capacity'))" ) ) pending = await session.scalar( select(func.count()) .select_from(SafetyTask) .where(SafetyTask.status.in_([TaskStatus.pending, TaskStatus.processing])) ) if max_pending is not None and pending >= max_pending: raise QueueFull statement = ( insert(SafetyRequest) .values( message_id=record.message_id, request_fingerprint=record.request_fingerprint, processing_mode=record.processing_mode, config_version=record.config_version, verdict=record.verdict, task_id=record.task_id, rule_id=record.rule_id, reason_code=record.reason_code, rules_version=record.rules_version, created_at=record.created_at, purge_after=record.purge_after, ) .on_conflict_do_nothing(index_elements=["message_id"]) ) result = await session.execute(statement.returning(SafetyRequest.message_id)) created = result.scalar_one_or_none() is not None existing = await session.get(SafetyRequest, record.message_id, with_for_update=True) assert existing if existing.request_fingerprint != record.request_fingerprint: raise ConflictError if created and task is not None: session.add(task) return existing, created async def add_task(self, task: SafetyTask) -> SafetyTask: async with self.sessions.begin() as session: session.add(task) return task async def task(self, task_id: uuid.UUID) -> SafetyTask | None: async with self.sessions.begin() as session: task = await session.get(SafetyTask, task_id, with_for_update=True) if ( task and task.status in {TaskStatus.pending, TaskStatus.processing} and task.expires_at <= datetime.now(UTC) ): task.status = TaskStatus.failed task.finished_at = datetime.now(UTC) task.purge_after = task.finished_at + timedelta(days=30) return task async def claim(self, owner: str) -> SafetyTask | None: async with self.sessions.begin() as session: row = ( await session.execute( text( """ WITH candidate AS ( SELECT t.id, (c.config->'task'->>'lease_sec')::integer AS lease_sec FROM message_safety.safety_tasks t JOIN message_safety.config_versions c ON c.version=t.config_version WHERE (t.status='pending' AND COALESCE(t.next_attempt_at, now()) <= now()) OR (t.status='processing' AND t.lease_until < now()) ORDER BY t.created_at FOR UPDATE OF t SKIP LOCKED LIMIT 1 ) UPDATE message_safety.safety_tasks t SET status='processing', lease_owner=:owner, lease_until=now() + make_interval(secs => candidate.lease_sec), lease_generation=lease_generation+1, attempt_count=attempt_count+1, updated_at=now() FROM candidate WHERE t.id=candidate.id RETURNING t.id """ ), {"owner": owner}, ) ).scalar_one_or_none() return await session.get(SafetyTask, row) if row else None async def heartbeat( self, task_id: uuid.UUID, owner: str, generation: int, lease_sec: int ) -> bool: async with self.sessions.begin() as session: result = await session.execute( update(SafetyTask) .where( SafetyTask.id == task_id, SafetyTask.status == TaskStatus.processing, SafetyTask.lease_owner == owner, SafetyTask.lease_generation == generation, ) .values(lease_until=func.now() + text(f"interval '{int(lease_sec)} seconds'")) ) return result.rowcount == 1 async def finish( self, task_id: uuid.UUID, owner: str, generation: int, *, allow: bool, rule_id: str, signatures_version: str | None = None, ) -> bool: now = datetime.now(UTC) values = { "status": TaskStatus.allowed if allow else TaskStatus.denied, "verdict": "allow" if allow else "deny", "rule_id": rule_id, "reason_code": None if allow else "message_blocked", "finished_at": now, "purge_after": now + timedelta(days=30), "updated_at": now, "lease_owner": None, "lease_until": None, } if signatures_version is not None: values["signatures_version"] = signatures_version async with self.sessions.begin() as session: result = await session.execute( update(SafetyTask) .where( SafetyTask.id == task_id, SafetyTask.status == TaskStatus.processing, SafetyTask.lease_owner == owner, SafetyTask.lease_generation == generation, SafetyTask.lease_until > func.now(), ) .values(**values) ) if result.rowcount == 1: await session.execute( update(SafetyRequest) .where(SafetyRequest.task_id == task_id) .values( verdict="allow" if allow else "deny", rule_id=rule_id, reason_code=None if allow else "message_blocked", ) ) return result.rowcount == 1 async def retry_or_fail(self, task: SafetyTask, max_attempts: int) -> None: async with self.sessions.begin() as session: terminal = task.attempt_count >= max_attempts or task.expires_at <= datetime.now(UTC) await session.execute( update(SafetyTask) .where( SafetyTask.id == task.id, SafetyTask.lease_owner == task.lease_owner, SafetyTask.lease_generation == task.lease_generation, ) .values( status=TaskStatus.failed if terminal else TaskStatus.pending, lease_owner=None, lease_until=None, next_attempt_at=None if terminal else datetime.now(UTC) + timedelta(seconds=2**task.attempt_count), finished_at=datetime.now(UTC) if terminal else None, ) ) async def audit(self, event: SafetyAudit) -> None: async with self.sessions.begin() as session: session.add(event)