from __future__ import annotations import os import ssl import uuid from datetime import datetime from enum import StrEnum from typing import Any from sqlalchemy import ( BigInteger, CheckConstraint, DateTime, Enum, ForeignKey, Index, Integer, LargeBinary, String, Text, UniqueConstraint, text, ) from sqlalchemy.dialects.postgresql import ARRAY, JSONB, UUID from sqlalchemy.ext.asyncio import AsyncAttrs, AsyncEngine, async_sessionmaker, create_async_engine from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column SCHEMA = "message_safety" def postgres_ssl_context() -> ssl.SSLContext: ca_file = os.environ.get("PG_CA_FILE") if not ca_file: raise RuntimeError("PG_CA_FILE is required") context = ssl.create_default_context(cafile=ca_file) context.check_hostname = True context.verify_mode = ssl.CERT_REQUIRED return context class Base(AsyncAttrs, DeclarativeBase): pass class TaskStatus(StrEnum): pending = "pending" processing = "processing" allowed = "allowed" denied = "denied" failed = "failed" class SafetyRequest(Base): __tablename__ = "safety_requests" __table_args__ = ( CheckConstraint("octet_length(request_fingerprint)=32", name="ck_request_fingerprint"), CheckConstraint("verdict IN ('allow','deny','pending')", name="ck_request_verdict"), CheckConstraint("processing_mode IN ('standard','mock')", name="ck_request_mode"), Index("ix_request_purge", "purge_after"), {"schema": SCHEMA}, ) message_id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True) request_fingerprint: Mapped[bytes] = mapped_column(LargeBinary(32)) processing_mode: Mapped[str] = mapped_column(String(16)) config_version: Mapped[int] = mapped_column( BigInteger, ForeignKey(f"{SCHEMA}.config_versions.version", ondelete="RESTRICT") ) verdict: Mapped[str] = mapped_column(String(8)) task_id: Mapped[uuid.UUID | None] = mapped_column(UUID(as_uuid=True)) rule_id: Mapped[str | None] = mapped_column(String(128)) reason_code: Mapped[str | None] = mapped_column(String(64)) rules_version: Mapped[str] = mapped_column(String(128)) created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True)) purge_after: Mapped[datetime] = mapped_column(DateTime(timezone=True)) class ConfigVersion(Base): __tablename__ = "config_versions" __table_args__ = ( CheckConstraint("state IN ('draft','active','retired')", name="ck_config_state"), Index( "uq_config_one_active", "state", unique=True, postgresql_where=text("state = 'active'") ), {"schema": SCHEMA}, ) id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4) version: Mapped[int] = mapped_column(BigInteger, unique=True) schema_version: Mapped[int] = mapped_column(Integer) state: Mapped[str] = mapped_column(String(16)) config: Mapped[dict[str, Any]] = mapped_column(JSONB) config_sha256: Mapped[bytes] = mapped_column(LargeBinary(32)) created_by: Mapped[str] = mapped_column(String(128)) created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True)) approved_by: Mapped[str | None] = mapped_column(String(128)) approved_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) activated_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) retired_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) class SafetyTask(Base): __tablename__ = "safety_tasks" __table_args__ = ( CheckConstraint("octet_length(request_fingerprint)=32", name="ck_task_fingerprint"), CheckConstraint("octet_length(content_sha256)=32", name="ck_task_sha"), CheckConstraint("processing_mode='standard'", name="ck_task_standard"), CheckConstraint("attempt_count>=0 AND lease_generation>=0", name="ck_task_counts"), CheckConstraint( "(status='allowed' AND verdict='allow') OR " "(status='denied' AND verdict='deny' AND reason_code='message_blocked') OR " "(status='failed' AND verdict IS NULL) OR " "(status IN ('pending','processing') AND verdict IS NULL)", name="ck_task_terminal", ), Index("ix_task_queue", "status", "next_attempt_at", "created_at"), Index("ix_task_lease", "status", "lease_until"), Index("ix_task_retention", "finished_at"), {"schema": SCHEMA}, ) id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4) message_id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), unique=True) attachment_id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True)) request_fingerprint: Mapped[bytes] = mapped_column(LargeBinary(32)) content_sha256: Mapped[bytes] = mapped_column(LargeBinary(32)) processing_mode: Mapped[str] = mapped_column(String(16), default="standard") config_version: Mapped[int] = mapped_column( BigInteger, ForeignKey(f"{SCHEMA}.config_versions.version", ondelete="RESTRICT") ) status: Mapped[TaskStatus] = mapped_column(Enum(TaskStatus, name="task_status", schema=SCHEMA)) attempt_count: Mapped[int] = mapped_column(Integer, default=0) lease_generation: Mapped[int] = mapped_column(Integer, default=0) lease_owner: Mapped[str | None] = mapped_column(String(128)) lease_until: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) next_attempt_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True)) quarantine_object_key: Mapped[str] = mapped_column(Text) quarantine_version_id: Mapped[str] = mapped_column(String(512)) quarantine_etag: Mapped[str] = mapped_column(String(512)) declared_mime: Mapped[str] = mapped_column(String(127)) declared_size_bytes: Mapped[int] = mapped_column(BigInteger) declared_checksum: Mapped[str] = mapped_column(String(71)) verdict: Mapped[str | None] = mapped_column(String(8)) rule_id: Mapped[str | None] = mapped_column(String(128)) reason_code: Mapped[str | None] = mapped_column(String(64)) rules_version: Mapped[str] = mapped_column(String(128)) detector_version: Mapped[str] = mapped_column(String(128)) scanner_engine: Mapped[str] = mapped_column(String(32)) signatures_version: Mapped[str] = mapped_column(String(128)) origin_trace_id: Mapped[bytes | None] = mapped_column(LargeBinary(16)) origin_span_id: Mapped[bytes | None] = mapped_column(LargeBinary(8)) origin_trace_flags: Mapped[int | None] = mapped_column(Integer) origin_tracestate: Mapped[str | None] = mapped_column(String(512)) created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True)) updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True)) finished_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) purge_after: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) class FileVerdictCache(Base): __tablename__ = "file_verdict_cache" __table_args__ = ( UniqueConstraint( "content_sha256", "config_version", "rules_version", "detector_version", "scanner_engine", "signatures_version", name="uq_file_cache_key", ), CheckConstraint("verdict IN ('allow','deny')", name="ck_file_cache_verdict"), CheckConstraint("octet_length(content_sha256)=32", name="ck_file_cache_sha"), Index("ix_file_cache_expiry", "expires_at"), {"schema": SCHEMA}, ) id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4) content_sha256: Mapped[bytes] = mapped_column(LargeBinary(32)) config_version: Mapped[int] = mapped_column( BigInteger, ForeignKey(f"{SCHEMA}.config_versions.version", ondelete="RESTRICT") ) rules_version: Mapped[str] = mapped_column(String(128)) detector_version: Mapped[str] = mapped_column(String(128)) scanner_engine: Mapped[str] = mapped_column(String(32)) signatures_version: Mapped[str] = mapped_column(String(128)) verdict: Mapped[str] = mapped_column(String(8)) rule_id: Mapped[str] = mapped_column(String(128)) reason_code: Mapped[str | None] = mapped_column(String(64)) created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True)) expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True)) class TextRulesCache(Base): __tablename__ = "text_rules_cache" __table_args__ = ( UniqueConstraint("analysis_sha256", "rules_version", name="uq_text_cache_key"), CheckConstraint("result IN ('allow','deny')", name="ck_text_cache_result"), CheckConstraint("octet_length(analysis_sha256)=32", name="ck_text_cache_sha"), Index("ix_text_cache_expiry", "expires_at"), {"schema": SCHEMA}, ) id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4) analysis_sha256: Mapped[bytes] = mapped_column(LargeBinary(32)) rules_version: Mapped[str] = mapped_column(String(128)) result: Mapped[str] = mapped_column(String(8)) deny_rule_id: Mapped[str | None] = mapped_column(String(128)) monitor_rule_ids: Mapped[list[str]] = mapped_column(ARRAY(String(128)), default=list) normalization_flags: Mapped[list[str]] = mapped_column(ARRAY(String(32)), default=list) created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True)) expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True)) class LinkVerdictCache(Base): __tablename__ = "link_verdict_cache" __table_args__ = ( UniqueConstraint( "canonical_url_sha256", "rules_version", "config_version", name="uq_link_key" ), CheckConstraint("verdict IN ('allow','deny')", name="ck_link_verdict"), Index("ix_link_cache_expiry", "expires_at"), {"schema": SCHEMA}, ) id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4) canonical_url_sha256: Mapped[bytes] = mapped_column(LargeBinary(32)) rules_version: Mapped[str] = mapped_column(String(128)) config_version: Mapped[int] = mapped_column( BigInteger, ForeignKey(f"{SCHEMA}.config_versions.version", ondelete="RESTRICT") ) verdict: Mapped[str] = mapped_column(String(8)) rule_id: Mapped[str | None] = mapped_column(String(128)) reason_code: Mapped[str | None] = mapped_column(String(64)) first_seen_at: Mapped[datetime] = mapped_column(DateTime(timezone=True)) last_seen_at: Mapped[datetime] = mapped_column(DateTime(timezone=True)) hit_count: Mapped[int] = mapped_column(BigInteger, default=1) expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True)) class SafetyAudit(Base): __tablename__ = "safety_audit" __table_args__ = ( CheckConstraint( "event IN ('received','task_created','rule_hit','rule_hit_monitor'," "'scan_completed','dependency_failed','mock_forced_allow'," "'mock_forced_deny','config_activated')", name="ck_audit_event", ), CheckConstraint("processing_mode IN ('standard','mock')", name="ck_audit_mode"), Index("ix_audit_purge", "purge_after"), Index("ix_audit_message", "message_id", "created_at"), {"schema": SCHEMA}, ) id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4) request_id: Mapped[str | None] = mapped_column(String(64)) message_id: Mapped[uuid.UUID | None] = mapped_column(UUID(as_uuid=True)) task_id: Mapped[uuid.UUID | None] = mapped_column(UUID(as_uuid=True)) event: Mapped[str] = mapped_column(String(32)) processing_mode: Mapped[str] = mapped_column(String(16)) config_version: Mapped[int] = mapped_column( BigInteger, ForeignKey(f"{SCHEMA}.config_versions.version", ondelete="RESTRICT") ) verdict: Mapped[str | None] = mapped_column(String(8)) rule_id: Mapped[str | None] = mapped_column(String(128)) rules_version: Mapped[str | None] = mapped_column(String(128)) normalization_flags: Mapped[list[str]] = mapped_column(ARRAY(String(32)), default=list) duration_ms: Mapped[int | None] = mapped_column(Integer) error_category: Mapped[str | None] = mapped_column(String(64)) created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True)) purge_after: Mapped[datetime] = mapped_column(DateTime(timezone=True)) def engine_and_sessions(url: str) -> tuple[AsyncEngine, async_sessionmaker]: engine = create_async_engine( url, pool_pre_ping=True, connect_args={"ssl": postgres_ssl_context()}, ) return engine, async_sessionmaker(engine, expire_on_commit=False)