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