Добавлен сбор телеметрии на ВМ2

This commit is contained in:
mi
2026-08-26 15:12:51 +03:00
parent 728b9826a3
commit f989097484
32 changed files with 2029 additions and 134 deletions
@@ -8,11 +8,14 @@ import boto3
import dns.asyncresolver
from botocore.config import Config
from botocore.exceptions import BotoCoreError, ClientError
from opentelemetry import trace
from app.contracts import Attachment
from app.file_pipeline import DependencyFailure, ObjectChanged
from app.url_policy import DnsError, DnsNxDomain
tracer = trace.get_tracer("message-safety.dependencies")
class TrustedDnsResolver:
def __init__(self, nameservers: list[str]) -> None:
@@ -22,6 +25,14 @@ class TrustedDnsResolver:
async def resolve(
self, hostname: str
) -> tuple[ipaddress.IPv4Address | ipaddress.IPv6Address, ...]:
with tracer.start_as_current_span(
"message_safety.dns.resolve", attributes={"server.address.type": "domain"}
):
return await self._resolve(hostname)
async def _resolve(
self, hostname: str
) -> tuple[ipaddress.IPv4Address | ipaddress.IPv6Address, ...]:
found: list[ipaddress.IPv4Address | ipaddress.IPv6Address] = []
try:
@@ -3,6 +3,7 @@ from __future__ import annotations
import hmac
import json
import uuid
from time import perf_counter
from typing import Annotated
from fastapi import Depends, FastAPI, Header, Request
@@ -14,6 +15,7 @@ from app.db import TaskStatus
from app.rate_limit import RateLimited
from app.repository import ConflictError
from app.service import CapabilityUnavailable, SafetyService, TaskFailed
from app.telemetry import logger, record_poll
CHECK_ADAPTER = TypeAdapter(CheckRequest)
MAX_BODY = 16_384
@@ -53,14 +55,39 @@ def create_app(service: SafetyService, token: str) -> FastAPI:
@app.middleware("http")
async def request_context(request: Request, call_next):
started = perf_counter()
supplied = request.headers.get("X-Request-ID")
try:
request.state.request_id = str(uuid.UUID(supplied)) if supplied else str(uuid.uuid4())
except ValueError:
request.state.request_id = str(uuid.uuid4())
response = await call_next(request)
try:
response = await call_next(request)
except Exception:
logger.error(
"Safety request failed",
extra={
"event": "http_request_failed",
"request_id": request.state.request_id,
"route": _route_template(request),
"method": request.method,
},
)
raise
response.headers["X-Request-ID"] = request.state.request_id
response.headers["Cache-Control"] = "no-store"
if request.url.path != "/health/live":
logger.info(
"Safety request completed",
extra={
"event": "http_request_completed",
"request_id": request.state.request_id,
"route": _route_template(request),
"method": request.method,
"status_code": response.status_code,
"duration_ms": round((perf_counter() - started) * 1000, 3),
},
)
return response
@app.get("/health/live")
@@ -197,8 +224,10 @@ def create_app(service: SafetyService, token: str) -> FastAPI:
)
task = await service.repository.task(parsed)
if not task:
record_poll("not_found")
return error(404, "task_not_found", request.state.request_id)
if task.status == TaskStatus.failed:
record_poll("failed")
return error(
503,
"task_failed",
@@ -211,6 +240,7 @@ def create_app(service: SafetyService, token: str) -> FastAPI:
},
)
if task.status in {TaskStatus.pending, TaskStatus.processing}:
record_poll(task.status.value)
result: Pending | Verdict = Pending(
config_version=task.config_version,
task_id=task.id,
@@ -218,6 +248,7 @@ def create_app(service: SafetyService, token: str) -> FastAPI:
rules_version=task.rules_version,
)
else:
record_poll("terminal")
result = service._verdict(
task.status == TaskStatus.allowed,
task.processing_mode,
@@ -233,3 +264,8 @@ def create_app(service: SafetyService, token: str) -> FastAPI:
return response
return app
def _route_template(request: Request) -> str:
route = request.scope.get("route")
return getattr(route, "path", "unmatched")
@@ -145,6 +145,10 @@ class SafetyTask(Base):
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))
@@ -10,12 +10,14 @@ from dataclasses import dataclass
from pathlib import Path
from typing import Protocol
from opentelemetry import trace
from PIL import Image, UnidentifiedImageError
from pillow_heif import register_heif_opener
from app.contracts import Attachment
register_heif_opener()
tracer = trace.get_tracer("message-safety.dependencies")
class ObjectChanged(RuntimeError):
@@ -122,6 +124,13 @@ class ClamAvInstream:
self.host, self.port, self.timeout = host, port, timeout
async def scan(self, chunks: AsyncIterator[bytes]) -> str | None:
with tracer.start_as_current_span(
"message_safety.clamav.scan",
attributes={"server.address.type": "clamav"},
):
return await self._scan(chunks)
async def _scan(self, chunks: AsyncIterator[bytes]) -> str | None:
async def operation() -> str | None:
reader, writer = await asyncio.open_connection(self.host, self.port)
try:
@@ -148,6 +157,10 @@ class ClamAvInstream:
raise DependencyFailure("ClamAV unavailable") from exc
async def signatures_version(self) -> str:
with tracer.start_as_current_span("message_safety.clamav.version"):
return await self._signatures_version()
async def _signatures_version(self) -> str:
try:
reader, writer = await asyncio.wait_for(
asyncio.open_connection(self.host, self.port), 2.0
@@ -12,6 +12,7 @@ from app.file_pipeline import ClamAvInstream, DependencyFailure
from app.repository import Repository
from app.service import SafetyService
from app.settings import BootstrapSettings, EmergencyMode
from app.telemetry import init_telemetry, instrument_fastapi, shutdown_telemetry
async def build_runtime() -> tuple[object, object]:
@@ -48,19 +49,24 @@ async def build_runtime() -> tuple[object, object]:
signatures_version=signatures_version,
)
app = create_app(service, settings.service_token.get_secret_value())
instrument_fastapi(app)
return app, engine
async def serve() -> None:
settings = BootstrapSettings()
app, engine = await build_runtime()
init_telemetry("api")
engine = None
try:
app, engine = await build_runtime()
server = uvicorn.Server(
uvicorn.Config(app, host=settings.host, port=settings.port, proxy_headers=False)
)
await server.serve()
finally:
await engine.dispose()
if engine is not None:
await engine.dispose()
shutdown_telemetry()
def run() -> None:
@@ -2,6 +2,7 @@ from __future__ import annotations
import uuid
from datetime import UTC, datetime, timedelta
from time import perf_counter
from app.config import ActiveConfig
from app.contracts import CheckRequest, FileCheck, Pending, TextCheck, Verdict
@@ -19,6 +20,14 @@ 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.telemetry import (
current_trace_fields,
record_cache,
record_check,
record_dependency,
record_file_rejection,
record_runtime_state,
)
from app.url_policy import DnsError, Resolver, canonicalize, check_url, extract_urls
@@ -54,6 +63,7 @@ class SafetyService:
self.rate_limiter = ConservativeRateLimiter(
config.document["rate"]["text_rps"], config.document["rate"]["file_rps"]
)
record_runtime_state("mock" if mode.mock else "standard", config.version)
def _verdict(
self,
@@ -74,6 +84,22 @@ class SafetyService:
)
async def check(self, request: CheckRequest) -> Verdict | Pending:
started = perf_counter()
outcome = "error"
try:
result = await self._check(request)
outcome = "pending" if isinstance(result, Pending) else result.verdict
return result
finally:
record_check(
request.content_kind,
"mock" if self.mode.mock else "standard",
outcome,
(perf_counter() - started) * 1000,
self.config.version,
)
async def _check(self, request: CheckRequest) -> Verdict | Pending:
digest = fingerprint(request)
existing = await self.repository.get_request(request.message_id)
if existing:
@@ -101,6 +127,7 @@ class SafetyService:
cache = await self.repository.text_cache(
normalized.analysis_sha256, self.config.rules_version
)
record_cache("text_rules", cache is not None)
if cache:
deny_rule = cache.deny_rule_id
else:
@@ -145,6 +172,7 @@ class SafetyService:
cached_link = await self.repository.link_cache(
canonical.digest, self.config.rules_version, self.config.version
)
record_cache("link_verdict", cached_link is not None)
if cached_link and cached_link.verdict == "deny":
return await self._persist_sync(
request,
@@ -160,7 +188,9 @@ class SafetyService:
_, rule = await check_url(
raw, self.resolver, self.config.document["link"]["dns_lookup_timeout_sec"]
)
record_dependency("dns", "resolve", "success")
except DnsError as exc:
record_dependency("dns", "resolve", "error")
raise CapabilityUnavailable("dns") from exc
if rule != "url.nxdomain" and not cached_link:
now = datetime.now(UTC)
@@ -200,6 +230,7 @@ class SafetyService:
set(self.config.document["file_policy"]["enabled_mime_types"]),
)
if rule:
record_file_rejection(rule)
return await self._persist_sync(
request, digest, self._verdict(False, "standard", rule, self.config.rules_version)
)
@@ -207,6 +238,7 @@ class SafetyService:
cached = await self.repository.file_cache(
content_digest, self.config, self.signatures_version
)
record_cache("file_verdict", cached is not None)
if cached:
return await self._persist_sync(
request,
@@ -220,6 +252,7 @@ class SafetyService:
)
now = datetime.now(UTC)
task_id = uuid.uuid4()
trace_id, span_id, trace_flags, tracestate = current_trace_fields()
task = SafetyTask(
id=task_id,
message_id=request.message_id,
@@ -243,6 +276,10 @@ class SafetyService:
detector_version=self.config.detector.version,
scanner_engine="clamav",
signatures_version=self.signatures_version,
origin_trace_id=trace_id,
origin_span_id=span_id,
origin_trace_flags=trace_flags,
origin_tracestate=tracestate,
created_at=now,
updated_at=now,
)
@@ -262,6 +299,7 @@ class SafetyService:
row, task, max_pending=self.config.document["task"]["max_pending"]
)
except QueueFull as exc:
record_dependency("postgresql", "queue_capacity", "full")
raise CapabilityUnavailable("queue_capacity") from exc
if created:
await self._audit(
@@ -0,0 +1,377 @@
from __future__ import annotations
import json
import logging
import os
import re
import sys
from collections.abc import Mapping, MutableMapping
from dataclasses import dataclass
from datetime import UTC, datetime
from typing import Any
from fastapi import FastAPI
from opentelemetry import metrics, trace
from opentelemetry.exporter.otlp.proto.grpc._log_exporter import OTLPLogExporter
from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import OTLPMetricExporter
from opentelemetry.exporter.otlp.proto.grpc.trace_exporter import OTLPSpanExporter
from opentelemetry.instrumentation.botocore import BotocoreInstrumentor
from opentelemetry.instrumentation.fastapi import FastAPIInstrumentor
from opentelemetry.instrumentation.httpx import HTTPXClientInstrumentor
from opentelemetry.instrumentation.redis import RedisInstrumentor
from opentelemetry.instrumentation.sqlalchemy import SQLAlchemyInstrumentor
from opentelemetry.propagate import set_global_textmap
from opentelemetry.sdk._logs import LoggerProvider, LoggingHandler
from opentelemetry.sdk._logs.export import BatchLogRecordProcessor
from opentelemetry.sdk.metrics import MeterProvider
from opentelemetry.sdk.metrics.export import PeriodicExportingMetricReader
from opentelemetry.sdk.resources import Resource
from opentelemetry.sdk.trace import TracerProvider
from opentelemetry.sdk.trace.export import BatchSpanProcessor
from opentelemetry.trace import Link, SpanContext, TraceFlags, TraceState
from opentelemetry.trace.propagation.tracecontext import TraceContextTextMapPropagator
SERVICE_NAME = "message-safety"
SERVICE_NAMESPACE = "han-chat"
_SENSITIVE_KEY = re.compile(
r"(authorization|cookie|token|secret|password|phone|email|message|text|filename|"
r"object[_-]?key|url|dsn|statement)",
re.IGNORECASE,
)
_SENSITIVE_VALUE = re.compile(
r"(?:https?://\S+\?\S+)|(?:[\w.+-]+@[\w.-]+\.[A-Za-z]{2,})|"
r"(?:\+?\d[\d ()-]{8,}\d)|(?:bearer\s+\S+)",
re.IGNORECASE,
)
_STANDARD_LOG_RECORD = set(logging.makeLogRecord({}).__dict__)
_SAFE_STRUCTURED_KEYS = {
"deployment.environment",
"duration_ms",
"environment",
"event",
"level",
"method",
"outcome",
"process_role",
"request_id",
"route",
"service.name",
"service.version",
"span_id",
"status_code",
"timestamp",
"trace_id",
}
_runtime: TelemetryRuntime | None = None
def _redact(value: Any, key: str = "") -> Any:
if key in _SAFE_STRUCTURED_KEYS:
return value
if _SENSITIVE_KEY.search(key):
return "[REDACTED]"
if isinstance(value, dict):
return {str(k)[:64]: _redact(v, str(k)) for k, v in list(value.items())[:64]}
if isinstance(value, (list, tuple)):
return [_redact(item) for item in value[:32]]
if isinstance(value, str):
return _SENSITIVE_VALUE.sub("[REDACTED]", value)
return value
class RedactionFilter(logging.Filter):
def filter(self, record: logging.LogRecord) -> bool:
record.msg = _redact(record.getMessage())
record.args = ()
for key in tuple(record.__dict__):
if key not in _STANDARD_LOG_RECORD:
record.__dict__[key] = _redact(record.__dict__[key], key)
return True
def _redact_span_attributes(attributes: MutableMapping[str, Any]) -> None:
for key in tuple(attributes):
key_text = str(key)
if _SENSITIVE_KEY.search(key_text) or key_text in {
"http.target",
"http.url",
"url.full",
"url.query",
"db.statement",
"db.query.text",
}:
attributes[key] = "[REDACTED]"
else:
attributes[key] = _redact(attributes[key], key_text)
class RedactingBatchSpanProcessor(BatchSpanProcessor):
"""Redact auto-instrumentation attributes before they enter the export queue."""
def on_end(self, span: Any) -> None:
attributes = getattr(span, "_attributes", None)
if isinstance(attributes, Mapping):
sanitized = dict(attributes)
_redact_span_attributes(sanitized)
span._attributes = sanitized
super().on_end(span)
class JsonFormatter(logging.Formatter):
def format(self, record: logging.LogRecord) -> str:
context = trace.get_current_span().get_span_context()
payload: dict[str, Any] = {
"timestamp": datetime.now(UTC)
.isoformat(timespec="milliseconds")
.replace("+00:00", "Z"),
"level": record.levelname,
"service.name": SERVICE_NAME,
"service.version": os.getenv("RELEASE_VERSION", "unknown"),
"environment": os.getenv("APP_ENV", "development"),
"event": getattr(record, "event", "application_log"),
"message": record.getMessage(),
}
if context.is_valid:
payload["trace_id"] = f"{context.trace_id:032x}"
payload["span_id"] = f"{context.span_id:016x}"
for key in (
"request_id",
"route",
"method",
"status_code",
"duration_ms",
"process_role",
"outcome",
):
value = getattr(record, key, None)
if value is not None:
payload[key] = value
return json.dumps(_redact(payload), ensure_ascii=False, separators=(",", ":"))
def configure_logging() -> logging.Logger:
logger = logging.getLogger("message_safety")
logger.setLevel(os.getenv("LOG_LEVEL", "INFO").upper())
logger.propagate = False
if not any(getattr(handler, "_message_safety_stdout", False) for handler in logger.handlers):
handler = logging.StreamHandler(sys.stdout)
handler._message_safety_stdout = True # type: ignore[attr-defined]
handler.addFilter(RedactionFilter())
handler.setFormatter(JsonFormatter())
logger.addHandler(handler)
return logger
logger = configure_logging()
@dataclass(slots=True)
class TelemetryRuntime:
tracer_provider: TracerProvider
meter_provider: MeterProvider
logger_provider: LoggerProvider
def shutdown(self) -> None:
for provider in (self.logger_provider, self.meter_provider, self.tracer_provider):
try:
provider.shutdown()
except Exception:
logger.error(
"Telemetry provider shutdown failed",
extra={"event": "telemetry_shutdown_failed"},
)
def _resource(role: str) -> Resource:
return Resource.create(
{
"service.name": SERVICE_NAME,
"service.namespace": SERVICE_NAMESPACE,
"service.version": os.getenv("RELEASE_VERSION", "unknown"),
"deployment.environment": os.getenv("APP_ENV", "development"),
"process.role": role,
}
)
def init_telemetry(role: str) -> TelemetryRuntime | None:
global _runtime
if _runtime is not None:
return _runtime
endpoint = os.getenv("OTEL_EXPORTER_OTLP_ENDPOINT", "").strip()
if not endpoint:
logger.info(
"OTLP endpoint is not configured; using JSON stdout",
extra={"event": "telemetry_stdout_only", "process_role": role},
)
return None
try:
resource = _resource(role)
insecure = endpoint.startswith("http://")
set_global_textmap(TraceContextTextMapPropagator())
tracer_provider = TracerProvider(resource=resource)
tracer_provider.add_span_processor(
RedactingBatchSpanProcessor(
OTLPSpanExporter(endpoint=endpoint, insecure=insecure, timeout=3),
max_queue_size=2048,
max_export_batch_size=512,
export_timeout_millis=3000,
)
)
trace.set_tracer_provider(tracer_provider)
metric_reader = PeriodicExportingMetricReader(
OTLPMetricExporter(endpoint=endpoint, insecure=insecure, timeout=3),
export_interval_millis=30_000,
export_timeout_millis=3000,
)
meter_provider = MeterProvider(resource=resource, metric_readers=[metric_reader])
metrics.set_meter_provider(meter_provider)
logger_provider = LoggerProvider(resource=resource)
logger_provider.add_log_record_processor(
BatchLogRecordProcessor(
OTLPLogExporter(endpoint=endpoint, insecure=insecure, timeout=3),
max_queue_size=2048,
max_export_batch_size=512,
export_timeout_millis=3000,
)
)
otlp_handler = LoggingHandler(level=logging.INFO, logger_provider=logger_provider)
otlp_handler.addFilter(RedactionFilter())
logger.addHandler(otlp_handler)
HTTPXClientInstrumentor().instrument()
SQLAlchemyInstrumentor().instrument(enable_commenter=False)
RedisInstrumentor().instrument()
BotocoreInstrumentor().instrument()
_runtime = TelemetryRuntime(tracer_provider, meter_provider, logger_provider)
logger.info(
"OpenTelemetry initialized",
extra={"event": "telemetry_initialized", "process_role": role},
)
return _runtime
except Exception:
logger.error(
"OpenTelemetry initialization failed; continuing with JSON stdout",
extra={"event": "telemetry_init_failed", "process_role": role},
)
return None
def instrument_fastapi(app: FastAPI) -> None:
if _runtime is None:
return
try:
FastAPIInstrumentor.instrument_app(app, excluded_urls="/health/live,/health/ready")
except Exception:
logger.error(
"FastAPI instrumentation failed; continuing",
extra={"event": "telemetry_instrumentation_failed"},
)
def shutdown_telemetry() -> None:
global _runtime
if _runtime is not None:
_runtime.shutdown()
_runtime = None
_meter = metrics.get_meter("message-safety.business")
checks = _meter.create_counter("message_safety.checks", description="Safety checks by outcome")
check_duration = _meter.create_histogram(
"message_safety.check.duration", unit="ms", description="Safety check duration"
)
cache_access = _meter.create_counter(
"message_safety.cache.access", description="Safety cache accesses"
)
worker_tasks = _meter.create_counter(
"message_safety.worker.tasks", description="Worker task outcomes"
)
worker_age = _meter.create_histogram(
"message_safety.worker.task.age", unit="s", description="Age of claimed worker tasks"
)
dependency_calls = _meter.create_counter(
"message_safety.dependency.calls", description="Dependency call outcomes"
)
polls = _meter.create_counter("message_safety.polls", description="Task poll outcomes")
file_rejections = _meter.create_counter(
"message_safety.file.rejections", description="File rejection categories"
)
runtime_state = _meter.create_up_down_counter(
"message_safety.runtime.state", description="Current runtime mode and config"
)
def record_check(
kind: str, mode: str, outcome: str, duration_ms: float, config_version: int
) -> None:
attributes = {
"content_kind": kind,
"processing_mode": mode,
"outcome": outcome,
"config_version": config_version,
}
checks.add(1, attributes)
check_duration.record(duration_ms, attributes)
def record_cache(cache: str, hit: bool) -> None:
cache_access.add(1, {"cache": cache, "result": "hit" if hit else "miss"})
def record_worker(outcome: str, age_seconds: float | None = None) -> None:
worker_tasks.add(1, {"outcome": outcome})
if age_seconds is not None:
worker_age.record(max(age_seconds, 0), {"outcome": outcome})
def record_dependency(dependency: str, operation: str, outcome: str) -> None:
dependency_calls.add(
1, {"dependency": dependency, "operation": operation, "outcome": outcome}
)
def record_poll(outcome: str) -> None:
polls.add(1, {"outcome": outcome})
def record_file_rejection(reason: str) -> None:
file_rejections.add(1, {"reason": reason})
def record_runtime_state(mode: str, config_version: int) -> None:
runtime_state.add(1, {"processing_mode": mode, "config_version": config_version})
def current_trace_fields() -> tuple[bytes | None, bytes | None, int | None, str | None]:
context = trace.get_current_span().get_span_context()
if not context.is_valid:
return None, None, None, None
return (
context.trace_id.to_bytes(16, "big"),
context.span_id.to_bytes(8, "big"),
int(context.trace_flags),
str(context.trace_state) or None,
)
def task_link(task: Any) -> Link | None:
trace_id = getattr(task, "origin_trace_id", None)
span_id = getattr(task, "origin_span_id", None)
if not trace_id or not span_id:
return None
try:
tracestate = getattr(task, "origin_tracestate", None)
context = SpanContext(
trace_id=int.from_bytes(trace_id, "big"),
span_id=int.from_bytes(span_id, "big"),
is_remote=True,
trace_flags=TraceFlags(getattr(task, "origin_trace_flags", 0) or 0),
trace_state=TraceState.from_header([tracestate] if tracestate else []),
)
return Link(context) if context.is_valid else None
except (TypeError, ValueError):
return None
@@ -5,6 +5,8 @@ import socket
import uuid
from datetime import UTC, datetime, timedelta
from opentelemetry import trace
from app.adapters import S3VersionReader
from app.config import validate_config
from app.contracts import Attachment
@@ -19,6 +21,15 @@ from app.file_pipeline import (
)
from app.repository import Repository
from app.settings import BootstrapSettings
from app.telemetry import (
init_telemetry,
record_dependency,
record_worker,
shutdown_telemetry,
task_link,
)
tracer = trace.get_tracer("message-safety.worker")
class Worker:
@@ -34,14 +45,30 @@ class Worker:
self.owner = f"{socket.gethostname()}:{uuid.uuid4()}"
async def once(self) -> bool:
task = await self.repository.claim(self.owner)
with tracer.start_as_current_span("message_safety.worker.claim"):
task = await self.repository.claim(self.owner)
if not task:
return False
link = task_link(task)
with tracer.start_as_current_span(
"message_safety.worker.process",
links=[link] if link else (),
attributes={
"messaging.operation.type": "process",
"message_safety.attempt": task.attempt_count,
},
):
await self._process(task)
return True
async def _process(self, task) -> None:
task_age = (datetime.now(UTC) - task.created_at).total_seconds()
row = await self.repository.config_version(task.config_version)
rules, detector, digest = validate_config(row.config, self.artifacts)
if digest != row.config_sha256:
await self.repository.retry_or_fail(task, row.config["task"]["max_attempts"])
return True
record_worker("config_error", task_age)
return
stop = asyncio.Event()
heartbeat = asyncio.create_task(
self._heartbeat(
@@ -65,18 +92,22 @@ class Worker:
body, _ = await collect_and_hash(
self.reader, attachment, max_size=row.config["file_policy"]["max_size_bytes"]
)
record_dependency("s3", "get_object", "success")
rule = detect_format(body, attachment.mime_type)
if not rule:
malware = await self.antivirus.scan(one_chunk(body))
record_dependency("clamav", "scan", "success")
rule = "file.malware_detected" if malware else None
finished = await self.repository.finish(
task.id,
self.owner,
task.lease_generation,
allow=rule is None,
rule_id=rule or "safety.all_checks_passed",
)
with tracer.start_as_current_span("message_safety.worker.finalize"):
finished = await self.repository.finish(
task.id,
self.owner,
task.lease_generation,
allow=rule is None,
rule_id=rule or "safety.all_checks_passed",
)
if finished:
record_worker("allow" if rule is None else "deny", task_age)
now = datetime.now(UTC)
await self.repository.put_file_cache(
FileVerdictCache(
@@ -111,19 +142,23 @@ class Worker:
)
)
except ObjectChanged:
await self.repository.finish(
task.id,
self.owner,
task.lease_generation,
allow=False,
rule_id="file.object_changed",
)
record_dependency("s3", "get_object", "object_changed")
with tracer.start_as_current_span("message_safety.worker.finalize"):
await self.repository.finish(
task.id,
self.owner,
task.lease_generation,
allow=False,
rule_id="file.object_changed",
)
record_worker("deny", task_age)
except DependencyFailure:
record_dependency("file_pipeline", "scan", "error")
await self.repository.retry_or_fail(task, row.config["task"]["max_attempts"])
record_worker("retry_or_fail", task_age)
finally:
stop.set()
await heartbeat
return True
async def _heartbeat(
self, task_id, generation: int, interval: int, lease: int, stop: asyncio.Event
@@ -144,25 +179,29 @@ class Worker:
async def serve() -> None:
settings = BootstrapSettings()
assert settings.database_url and settings.s3_access_key and settings.s3_secret_key
engine, sessions = engine_and_sessions(settings.database_url.get_secret_value())
worker = Worker(
Repository(sessions),
S3VersionReader(
settings.s3_endpoint_url,
settings.s3_bucket,
settings.s3_access_key.get_secret_value(),
settings.s3_secret_key.get_secret_value(),
),
ClamAvInstream(settings.clamav_host, settings.clamav_port),
settings.artifacts_dir,
)
init_telemetry("worker")
engine = None
try:
assert settings.database_url and settings.s3_access_key and settings.s3_secret_key
engine, sessions = engine_and_sessions(settings.database_url.get_secret_value())
worker = Worker(
Repository(sessions),
S3VersionReader(
settings.s3_endpoint_url,
settings.s3_bucket,
settings.s3_access_key.get_secret_value(),
settings.s3_secret_key.get_secret_value(),
),
ClamAvInstream(settings.clamav_host, settings.clamav_port),
settings.artifacts_dir,
)
async with asyncio.TaskGroup() as group:
for _ in range(settings.worker_concurrency):
group.create_task(worker.loop())
finally:
await engine.dispose()
if engine is not None:
await engine.dispose()
shutdown_telemetry()
def run() -> None: