Проект разделен на два репозитория
This commit is contained in:
@@ -0,0 +1,367 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from datetime import UTC, datetime, timedelta
|
||||
|
||||
from app.config import ActiveConfig
|
||||
from app.contracts import CheckRequest, FileCheck, Pending, TextCheck, Verdict
|
||||
from app.db import (
|
||||
LinkVerdictCache,
|
||||
SafetyAudit,
|
||||
SafetyRequest,
|
||||
SafetyTask,
|
||||
TaskStatus,
|
||||
TextRulesCache,
|
||||
)
|
||||
from app.file_pipeline import validate_metadata
|
||||
from app.fingerprint import fingerprint
|
||||
from app.normalization import normalize_text
|
||||
from app.rate_limit import ConservativeRateLimiter
|
||||
from app.repository import QueueFull, Repository
|
||||
from app.settings import EmergencyMode
|
||||
from app.url_policy import DnsError, Resolver, canonicalize, check_url, extract_urls
|
||||
|
||||
|
||||
class CapabilityUnavailable(RuntimeError):
|
||||
def __init__(self, category: str) -> None:
|
||||
self.category = category
|
||||
|
||||
|
||||
class TaskFailed(RuntimeError):
|
||||
def __init__(self, task_id: uuid.UUID) -> None:
|
||||
self.task_id = task_id
|
||||
|
||||
|
||||
class SafetyService:
|
||||
def __init__(
|
||||
self,
|
||||
repository: Repository,
|
||||
config: ActiveConfig,
|
||||
mode: EmergencyMode,
|
||||
resolver: Resolver,
|
||||
*,
|
||||
links_ready: bool = True,
|
||||
files_ready: bool = True,
|
||||
signatures_version: str = "unverified",
|
||||
) -> None:
|
||||
self.repository = repository
|
||||
self.config = config
|
||||
self.mode = mode
|
||||
self.resolver = resolver
|
||||
self.links_ready = links_ready
|
||||
self.files_ready = files_ready
|
||||
self.signatures_version = signatures_version
|
||||
self.rate_limiter = ConservativeRateLimiter(
|
||||
config.document["rate"]["text_rps"], config.document["rate"]["file_rps"]
|
||||
)
|
||||
|
||||
def _verdict(
|
||||
self,
|
||||
allow: bool,
|
||||
mode: str,
|
||||
rule: str,
|
||||
rules_version: str,
|
||||
*,
|
||||
config_version: int | None = None,
|
||||
) -> Verdict:
|
||||
return Verdict(
|
||||
verdict="allow" if allow else "deny",
|
||||
processing_mode=mode,
|
||||
config_version=self.config.version if config_version is None else config_version,
|
||||
rule_id=rule,
|
||||
reason_code=None if allow else "message_blocked",
|
||||
rules_version=rules_version,
|
||||
)
|
||||
|
||||
async def check(self, request: CheckRequest) -> Verdict | Pending:
|
||||
digest = fingerprint(request)
|
||||
existing = await self.repository.get_request(request.message_id)
|
||||
if existing:
|
||||
if existing.request_fingerprint != digest:
|
||||
from app.repository import ConflictError
|
||||
|
||||
raise ConflictError
|
||||
return await self._replay(existing)
|
||||
await self.rate_limiter.acquire(request.content_kind)
|
||||
if self.mode.mock:
|
||||
free = self.mode.text_free if request.content_kind == "text" else self.mode.file_free
|
||||
verdict = self._verdict(
|
||||
free,
|
||||
"mock",
|
||||
"safety.mock_forced_allow" if free else "safety.mock_forced_deny",
|
||||
"mock",
|
||||
)
|
||||
return await self._persist_sync(request, digest, verdict)
|
||||
if isinstance(request, TextCheck):
|
||||
return await self._check_text(request, digest)
|
||||
return await self._check_file(request, digest)
|
||||
|
||||
async def _check_text(self, request: TextCheck, digest: bytes) -> Verdict:
|
||||
normalized = normalize_text(request.text)
|
||||
cache = await self.repository.text_cache(
|
||||
normalized.analysis_sha256, self.config.rules_version
|
||||
)
|
||||
if cache:
|
||||
deny_rule = cache.deny_rule_id
|
||||
else:
|
||||
result = self.config.rules.evaluate(normalized.analysis)
|
||||
deny_rule = result.deny_rule
|
||||
now = datetime.now(UTC)
|
||||
await self.repository.put_text_cache(
|
||||
TextRulesCache(
|
||||
analysis_sha256=normalized.analysis_sha256,
|
||||
rules_version=self.config.rules_version,
|
||||
result="deny" if deny_rule else "allow",
|
||||
deny_rule_id=deny_rule,
|
||||
monitor_rule_ids=list(result.monitor_rules),
|
||||
normalization_flags=list(normalized.flags),
|
||||
created_at=now,
|
||||
expires_at=now
|
||||
+ timedelta(seconds=self.config.document["cache"]["text_rule_ttl_sec"]),
|
||||
)
|
||||
)
|
||||
if deny_rule:
|
||||
return await self._persist_sync(
|
||||
request,
|
||||
digest,
|
||||
self._verdict(False, "standard", deny_rule, self.config.rules_version),
|
||||
)
|
||||
urls = extract_urls(
|
||||
normalized.analysis,
|
||||
maximum=self.config.document["link"]["max_per_message"],
|
||||
max_length=self.config.document["link"]["url_max_length"],
|
||||
)
|
||||
if urls and not self.links_ready:
|
||||
raise CapabilityUnavailable("dns")
|
||||
for raw in urls:
|
||||
try:
|
||||
canonical = canonicalize(raw)
|
||||
except PermissionError as exc:
|
||||
return await self._persist_sync(
|
||||
request,
|
||||
digest,
|
||||
self._verdict(False, "standard", str(exc), self.config.rules_version),
|
||||
)
|
||||
cached_link = await self.repository.link_cache(
|
||||
canonical.digest, self.config.rules_version, self.config.version
|
||||
)
|
||||
if cached_link and cached_link.verdict == "deny":
|
||||
return await self._persist_sync(
|
||||
request,
|
||||
digest,
|
||||
self._verdict(
|
||||
False,
|
||||
"standard",
|
||||
cached_link.rule_id or "url.malformed",
|
||||
self.config.rules_version,
|
||||
),
|
||||
)
|
||||
try:
|
||||
_, rule = await check_url(
|
||||
raw, self.resolver, self.config.document["link"]["dns_lookup_timeout_sec"]
|
||||
)
|
||||
except DnsError as exc:
|
||||
raise CapabilityUnavailable("dns") from exc
|
||||
if rule != "url.nxdomain" and not cached_link:
|
||||
now = datetime.now(UTC)
|
||||
await self.repository.put_link_cache(
|
||||
LinkVerdictCache(
|
||||
canonical_url_sha256=canonical.digest,
|
||||
rules_version=self.config.rules_version,
|
||||
config_version=self.config.version,
|
||||
verdict="deny" if rule else "allow",
|
||||
rule_id=rule,
|
||||
reason_code="message_blocked" if rule else None,
|
||||
first_seen_at=now,
|
||||
last_seen_at=now,
|
||||
hit_count=1,
|
||||
expires_at=now
|
||||
+ timedelta(seconds=self.config.document["cache"]["link_ttl_sec"]),
|
||||
)
|
||||
)
|
||||
if rule and rule != "url.nxdomain":
|
||||
return await self._persist_sync(
|
||||
request,
|
||||
digest,
|
||||
self._verdict(False, "standard", rule, self.config.rules_version),
|
||||
)
|
||||
return await self._persist_sync(
|
||||
request,
|
||||
digest,
|
||||
self._verdict(True, "standard", "safety.all_checks_passed", self.config.rules_version),
|
||||
)
|
||||
|
||||
async def _check_file(self, request: FileCheck, digest: bytes) -> Verdict | Pending:
|
||||
if not self.files_ready:
|
||||
raise CapabilityUnavailable("files")
|
||||
rule = validate_metadata(
|
||||
request.attachment,
|
||||
self.config.detector,
|
||||
set(self.config.document["file_policy"]["enabled_mime_types"]),
|
||||
)
|
||||
if rule:
|
||||
return await self._persist_sync(
|
||||
request, digest, self._verdict(False, "standard", rule, self.config.rules_version)
|
||||
)
|
||||
content_digest = bytes.fromhex(request.attachment.checksum[7:])
|
||||
cached = await self.repository.file_cache(
|
||||
content_digest, self.config, self.signatures_version
|
||||
)
|
||||
if cached:
|
||||
return await self._persist_sync(
|
||||
request,
|
||||
digest,
|
||||
self._verdict(
|
||||
cached.verdict == "allow",
|
||||
"standard",
|
||||
cached.rule_id,
|
||||
cached.rules_version,
|
||||
),
|
||||
)
|
||||
now = datetime.now(UTC)
|
||||
task_id = uuid.uuid4()
|
||||
task = SafetyTask(
|
||||
id=task_id,
|
||||
message_id=request.message_id,
|
||||
attachment_id=request.attachment.attachment_id,
|
||||
request_fingerprint=digest,
|
||||
content_sha256=content_digest,
|
||||
processing_mode="standard",
|
||||
config_version=self.config.version,
|
||||
status=TaskStatus.pending,
|
||||
attempt_count=0,
|
||||
lease_generation=0,
|
||||
expires_at=now
|
||||
+ timedelta(seconds=self.config.document["task"]["execution_deadline_sec"]),
|
||||
quarantine_object_key=request.attachment.quarantine_object_key,
|
||||
quarantine_version_id=request.attachment.quarantine_version_id,
|
||||
quarantine_etag=request.attachment.quarantine_etag,
|
||||
declared_mime=request.attachment.mime_type,
|
||||
declared_size_bytes=request.attachment.size_bytes,
|
||||
declared_checksum=request.attachment.checksum,
|
||||
rules_version=self.config.rules_version,
|
||||
detector_version=self.config.detector.version,
|
||||
scanner_engine="clamav",
|
||||
signatures_version=self.signatures_version,
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
row = SafetyRequest(
|
||||
message_id=request.message_id,
|
||||
request_fingerprint=digest,
|
||||
processing_mode="standard",
|
||||
config_version=self.config.version,
|
||||
verdict="pending",
|
||||
task_id=task_id,
|
||||
rules_version=self.config.rules_version,
|
||||
created_at=now,
|
||||
purge_after=now + timedelta(days=30),
|
||||
)
|
||||
try:
|
||||
stored, created = await self.repository.reserve_request(
|
||||
row, task, max_pending=self.config.document["task"]["max_pending"]
|
||||
)
|
||||
except QueueFull as exc:
|
||||
raise CapabilityUnavailable("queue_capacity") from exc
|
||||
if created:
|
||||
await self._audit(
|
||||
request.message_id,
|
||||
"task_created",
|
||||
"standard",
|
||||
"pending",
|
||||
None,
|
||||
task_id=task.id,
|
||||
)
|
||||
return await self._replay(stored)
|
||||
|
||||
async def _persist_sync(
|
||||
self, request: CheckRequest, digest: bytes, verdict: Verdict
|
||||
) -> Verdict:
|
||||
now = datetime.now(UTC)
|
||||
row = SafetyRequest(
|
||||
message_id=request.message_id,
|
||||
request_fingerprint=digest,
|
||||
processing_mode=verdict.processing_mode,
|
||||
config_version=verdict.config_version,
|
||||
verdict=verdict.verdict,
|
||||
rule_id=verdict.rule_id,
|
||||
reason_code=verdict.reason_code,
|
||||
rules_version=verdict.rules_version,
|
||||
created_at=now,
|
||||
purge_after=now + timedelta(days=30),
|
||||
)
|
||||
stored, created = await self.repository.reserve_request(row)
|
||||
if created:
|
||||
event = (
|
||||
f"mock_forced_{verdict.verdict}"
|
||||
if verdict.processing_mode == "mock"
|
||||
else ("rule_hit" if verdict.verdict == "deny" else "received")
|
||||
)
|
||||
await self._audit(
|
||||
request.message_id,
|
||||
event,
|
||||
verdict.processing_mode,
|
||||
verdict.verdict,
|
||||
verdict.rule_id,
|
||||
)
|
||||
replay = await self._replay(stored)
|
||||
assert isinstance(replay, Verdict)
|
||||
return replay
|
||||
|
||||
async def _replay(self, row: SafetyRequest) -> Verdict | Pending:
|
||||
if row.verdict == "pending":
|
||||
assert row.task_id
|
||||
task = await self.repository.task(row.task_id)
|
||||
if task and task.status in {TaskStatus.allowed, TaskStatus.denied}:
|
||||
return self._verdict(
|
||||
task.status == TaskStatus.allowed,
|
||||
task.processing_mode,
|
||||
task.rule_id or "safety.all_checks_passed",
|
||||
task.rules_version,
|
||||
config_version=task.config_version,
|
||||
)
|
||||
if task and task.status == TaskStatus.failed:
|
||||
raise TaskFailed(task.id)
|
||||
assert task
|
||||
return Pending(
|
||||
config_version=task.config_version,
|
||||
task_id=task.id,
|
||||
expires_at=task.expires_at,
|
||||
rules_version=task.rules_version,
|
||||
)
|
||||
return Verdict(
|
||||
verdict=row.verdict,
|
||||
processing_mode=row.processing_mode,
|
||||
config_version=row.config_version,
|
||||
rule_id=row.rule_id or "safety.all_checks_passed",
|
||||
reason_code=row.reason_code,
|
||||
rules_version=row.rules_version,
|
||||
)
|
||||
|
||||
async def _audit(
|
||||
self,
|
||||
message_id: uuid.UUID,
|
||||
event: str,
|
||||
mode: str,
|
||||
verdict: str,
|
||||
rule_id: str | None,
|
||||
*,
|
||||
task_id: uuid.UUID | None = None,
|
||||
) -> None:
|
||||
now = datetime.now(UTC)
|
||||
await self.repository.audit(
|
||||
SafetyAudit(
|
||||
id=uuid.uuid4(),
|
||||
message_id=message_id,
|
||||
task_id=task_id,
|
||||
event=event,
|
||||
processing_mode=mode,
|
||||
config_version=self.config.version,
|
||||
verdict=verdict,
|
||||
rule_id=rule_id,
|
||||
rules_version=self.config.rules_version if mode == "standard" else "mock",
|
||||
normalization_flags=[],
|
||||
created_at=now,
|
||||
purge_after=now + timedelta(days=self.config.document["retention"]["audit_days"]),
|
||||
)
|
||||
)
|
||||
Reference in New Issue
Block a user