Добавлен сбор телеметрии на ВМ2
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user