326 lines
13 KiB
Python
326 lines
13 KiB
Python
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 == "clamav",
|
|
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
|
|
) -> bool:
|
|
now = datetime.now(UTC)
|
|
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(
|
|
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 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)
|