#!/usr/bin/env python3
"""python_secure_logging_toolkit.py

A practical, vendor-neutral example of secure Python logging for production.

Features:
- Structured logging with stable event fields
- Redaction of sensitive keys and token-like values
- Safe exception logging helpers
- Basic validation of log level and field inputs
- Optional JSON output for machine consumption

This script is intended as a reference implementation and starting point.
It avoids hard-coded secrets and destructive actions.
"""

from __future__ import annotations

import argparse
import json
import logging
import os
import re
import sys
from dataclasses import dataclass, field
from typing import Any, Dict, Iterable, Optional


SENSITIVE_KEYS = {
    "password",
    "pass",
    "secret",
    "token",
    "access_token",
    "refresh_token",
    "authorization",
    "cookie",
    "session",
    "private_key",
    "connection_string",
    "api_key",
}

TOKEN_PATTERNS = [
    re.compile(r"(?i)bearer\s+[a-z0-9._-]+"),
    re.compile(r"(?i)api[_-]?key\s*[:=]\s*[^\s,;]+"),
    re.compile(r"(?i)secret\s*[:=]\s*[^\s,;]+"),
]

VALID_LEVELS = {"DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL"}


@dataclass
class LogEvent:
    """Structured log event."""

    message: str
    level: str = "INFO"
    event_type: str = "app.event"
    service: str = "example-service"
    environment: str = "production"
    request_id: Optional[str] = None
    user_id: Optional[str] = None
    extra_fields: Dict[str, Any] = field(default_factory=dict)


def validate_level(level: str) -> str:
    normalized = level.upper().strip()
    if normalized not in VALID_LEVELS:
        raise ValueError(f"Invalid log level: {level!r}. Choose from: {sorted(VALID_LEVELS)}")
    return normalized


def sanitize_scalar(value: Any) -> Any:
    """Redact token-like strings while preserving useful context."""
    if value is None:
        return None
    if isinstance(value, (int, float, bool)):
        return value
    text = str(value)
    for pattern in TOKEN_PATTERNS:
        text = pattern.sub("[REDACTED]", text)
    return text


def redact_mapping(data: Dict[str, Any]) -> Dict[str, Any]:
    redacted: Dict[str, Any] = {}
    for key, value in data.items():
        key_lower = str(key).lower()
        if key_lower in SENSITIVE_KEYS:
            redacted[key] = "[REDACTED]"
        elif isinstance(value, dict):
            redacted[key] = redact_mapping(value)
        elif isinstance(value, list):
            redacted[key] = [redact_mapping(v) if isinstance(v, dict) else sanitize_scalar(v) for v in value]
        else:
            redacted[key] = sanitize_scalar(value)
    return redacted


class RedactingJsonFormatter(logging.Formatter):
    def format(self, record: logging.LogRecord) -> str:
        payload: Dict[str, Any] = {
            "timestamp": self.formatTime(record, self.datefmt),
            "level": record.levelname,
            "logger": record.name,
            "message": sanitize_scalar(record.getMessage()),
        }

        for attr in ("event_type", "service", "environment", "request_id", "user_id"):
            value = getattr(record, attr, None)
            if value is not None:
                payload[attr] = sanitize_scalar(value)

        extra_fields = getattr(record, "extra_fields", None)
        if isinstance(extra_fields, dict) and extra_fields:
            payload["fields"] = redact_mapping(extra_fields)

        if record.exc_info:
            payload["exception"] = sanitize_scalar(self.formatException(record.exc_info))

        return json.dumps(payload, ensure_ascii=False, sort_keys=True)


class RedactingTextFormatter(logging.Formatter):
    def format(self, record: logging.LogRecord) -> str:
        base = f"{record.levelname} {record.name}: {sanitize_scalar(record.getMessage())}"
        extras = []

        for attr in ("event_type", "service", "environment", "request_id", "user_id"):
            value = getattr(record, attr, None)
            if value is not None:
                extras.append(f"{attr}={sanitize_scalar(value)}")

        extra_fields = getattr(record, "extra_fields", None)
        if isinstance(extra_fields, dict) and extra_fields:
            safe_fields = redact_mapping(extra_fields)
            extras.append(f"fields={safe_fields}")

        if record.exc_info:
            extras.append(f"exception={sanitize_scalar(self.formatException(record.exc_info))}")

        if extras:
            base = base + " | " + " ".join(extras)
        return base


def configure_logging(level: str, json_output: bool = False) -> logging.Logger:
    normalized_level = validate_level(level)
    logger = logging.getLogger("secure_logging_demo")
    logger.setLevel(getattr(logging, normalized_level))
    logger.propagate = False

    if logger.handlers:
        logger.handlers.clear()

    handler = logging.StreamHandler(sys.stdout)
    handler.setLevel(getattr(logging, normalized_level))
    formatter: logging.Formatter
    if json_output:
        formatter = RedactingJsonFormatter(datefmt="%Y-%m-%dT%H:%M:%S%z")
    else:
        formatter = RedactingTextFormatter(datefmt="%Y-%m-%dT%H:%M:%S%z")
    handler.setFormatter(formatter)
    logger.addHandler(handler)
    return logger


def log_event(logger: logging.Logger, event: LogEvent) -> None:
    extra = {
        "event_type": event.event_type,
        "service": event.service,
        "environment": event.environment,
        "request_id": event.request_id,
        "user_id": event.user_id,
        "extra_fields": event.extra_fields,
    }
    logger.log(getattr(logging, event.level), event.message, extra=extra)


def parse_key_value_items(items: Iterable[str]) -> Dict[str, str]:
    fields: Dict[str, str] = {}
    for item in items:
        if "=" not in item:
            raise ValueError(f"Invalid field {item!r}. Expected key=value.")
        key, value = item.split("=", 1)
        key = key.strip()
        if not key:
            raise ValueError(f"Invalid field {item!r}. Key cannot be empty.")
        fields[key] = value
    return fields


def demo(logger: logging.Logger, service: str, environment: str, request_id: Optional[str]) -> None:
    log_event(
        logger,
        LogEvent(
            message="startup complete",
            level="INFO",
            event_type="service.startup",
            service=service,
            environment=environment,
            request_id=request_id,
            extra_fields={"port": 8080, "region": os.getenv("APP_REGION", "unknown")},
        ),
    )

    log_event(
        logger,
        LogEvent(
            message="authentication failed",
            level="WARNING",
            event_type="security.auth.failure",
            service=service,
            environment=environment,
            request_id=request_id,
            extra_fields={
                "username": "alice@example.com",
                "password": "super-secret-value",
                "reason": "invalid credentials",
            },
        ),
    )

    try:
        raise RuntimeError("database query failed for user input: [REDACTED?]")
    except Exception:
        log_event(
            logger,
            LogEvent(
                message="request handling error",
                level="ERROR",
                event_type="app.request.error",
                service=service,
                environment=environment,
                request_id=request_id,
                extra_fields={"operation": "lookup", "status": "failed"},
            ),
        )
        logger.exception(
            "restricted diagnostic trace",
            extra={
                "event_type": "diagnostic.trace",
                "service": service,
                "environment": environment,
                "request_id": request_id,
                "extra_fields": {"note": "trace output should remain access-controlled"},
            },
        )


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description="Secure Python logging demo")
    parser.add_argument("--level", default="INFO", help="Log level: DEBUG, INFO, WARNING, ERROR, CRITICAL")
    parser.add_argument("--json", action="store_true", help="Emit JSON logs instead of text logs")
    parser.add_argument("--service", default="example-service", help="Stable service name to include in logs")
    parser.add_argument("--environment", default="production", help="Environment label such as dev, staging, or prod")
    parser.add_argument("--request-id", default=None, help="Optional correlation ID")
    parser.add_argument(
        "--field",
        action="append",
        default=[],
        metavar="KEY=VALUE",
        help="Optional structured field; repeatable",
    )
    return parser


def main() -> int:
    args = build_parser().parse_args()
    logger = configure_logging(args.level, json_output=args.json)
    extra_fields = parse_key_value_items(args.field)

    log_event(
        logger,
        LogEvent(
            message="secure logging initialized",
            level=validate_level(args.level),
            event_type="logging.init",
            service=args.service,
            environment=args.environment,
            request_id=args.request_id,
            extra_fields=redact_mapping(extra_fields),
        ),
    )

    demo(logger, args.service, args.environment, args.request_id)
    return 0


if __name__ == "__main__":
    raise SystemExit(main())