Files

378 lines
13 KiB
Python

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