package logger

import (
	"context"
	"log/slog"
	"regexp"
	"strings"
)

var (
	emailPattern = regexp.MustCompile(`(?i)[a-z0-9._%+\-]+@[a-z0-9.\-]+\.[a-z]{2,}`)
	panPattern   = regexp.MustCompile(`\b(?:\d[ -]?){13,19}\b`)
)

type RedactingHandler struct {
	next slog.Handler
}

func NewRedactingHandler(next slog.Handler) *RedactingHandler {
	return &RedactingHandler{next: next}
}

func (h *RedactingHandler) Enabled(ctx context.Context, level slog.Level) bool {
	return h.next.Enabled(ctx, level)
}

func (h *RedactingHandler) Handle(ctx context.Context, record slog.Record) error {
	redacted := slog.NewRecord(record.Time, record.Level, redactText(record.Message), record.PC)
	record.Attrs(func(attr slog.Attr) bool {
		redacted.AddAttrs(redactAttr(attr))
		return true
	})
	return h.next.Handle(ctx, redacted)
}

func (h *RedactingHandler) WithAttrs(attrs []slog.Attr) slog.Handler {
	redacted := make([]slog.Attr, 0, len(attrs))
	for _, attr := range attrs {
		redacted = append(redacted, redactAttr(attr))
	}
	return &RedactingHandler{next: h.next.WithAttrs(redacted)}
}

func (h *RedactingHandler) WithGroup(name string) slog.Handler {
	return &RedactingHandler{next: h.next.WithGroup(name)}
}

func redactAttr(attr slog.Attr) slog.Attr {
	attr.Value = redactValue(attr.Key, attr.Value)
	return attr
}

func redactValue(key string, value slog.Value) slog.Value {
	if sensitiveKey(key) {
		return slog.StringValue("[REDACTED]")
	}
	switch value.Kind() {
	case slog.KindString:
		return slog.StringValue(redactText(value.String()))
	case slog.KindGroup:
		attrs := value.Group()
		redacted := make([]slog.Attr, 0, len(attrs))
		for _, attr := range attrs {
			redacted = append(redacted, redactAttr(attr))
		}
		return slog.GroupValue(redacted...)
	default:
		return value
	}
}

func redactText(value string) string {
	value = emailPattern.ReplaceAllString(value, "[REDACTED_EMAIL]")
	value = panPattern.ReplaceAllStringFunc(value, func(candidate string) string {
		digits := strings.NewReplacer(" ", "", "-", "").Replace(candidate)
		if len(digits) >= 13 && len(digits) <= 19 && luhnValid(digits) {
			return "[REDACTED_CARD]"
		}
		return candidate
	})
	return value
}

func sensitiveKey(key string) bool {
	key = strings.ToLower(strings.TrimSpace(key))
	sensitiveFragments := []string{
		"password",
		"secret",
		"token",
		"authorization",
		"cookie",
		"cvv",
		"pan",
		"private_key",
		"api_key",
		"webhook",
		"mfa",
	}
	for _, fragment := range sensitiveFragments {
		if strings.Contains(key, fragment) {
			return true
		}
	}
	return false
}

func luhnValid(value string) bool {
	sum := 0
	double := false
	for i := len(value) - 1; i >= 0; i-- {
		digit := int(value[i] - '0')
		if digit < 0 || digit > 9 {
			return false
		}
		if double {
			digit *= 2
			if digit > 9 {
				digit -= 9
			}
		}
		sum += digit
		double = !double
	}
	return sum%10 == 0
}
