Реализованы сервисы ВМ2 - проверка сообщений и синхронизация с Б24 (деплой еще без перевода в боевой режим)
This commit is contained in:
@@ -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)
|
||||
Reference in New Issue
Block a user