Проект разделен на два репозитория

This commit is contained in:
mi
2026-08-14 15:42:45 +03:00
parent e06a77ee1d
commit bbef7a30c9
521 changed files with 2597 additions and 2302 deletions
@@ -0,0 +1,325 @@
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)