package correlation

import (
	"context"
	"net/http"
	"strings"

	"github.com/niels/banking-app/backend/internal/security"
)

const (
	HeaderRequestID     = "X-Request-ID"
	HeaderCorrelationID = "X-Correlation-ID"
	HeaderTraceParent   = "Traceparent"
)

type contextKey string

const valuesKey contextKey = "correlation_values"

type Values struct {
	RequestID string
	TraceID   string
	SpanID    string
}

func FromContext(ctx context.Context) Values {
	values, _ := ctx.Value(valuesKey).(Values)
	return values
}

func WithContext(ctx context.Context, values Values) context.Context {
	return context.WithValue(ctx, valuesKey, values)
}

func FromRequest(r *http.Request) Values {
	requestID := strings.TrimSpace(r.Header.Get(HeaderRequestID))
	if requestID == "" {
		requestID = strings.TrimSpace(r.Header.Get(HeaderCorrelationID))
	}
	if requestID == "" {
		requestID, _ = security.RandomHex(16)
	}

	traceID, _ := parseTraceParent(r.Header.Get(HeaderTraceParent))
	if traceID == "" {
		traceID, _ = security.RandomHex(16)
	}
	spanID, _ := security.RandomHex(8)

	return Values{
		RequestID: requestID,
		TraceID:   traceID,
		SpanID:    spanID,
	}
}

func ApplyResponseHeaders(w http.ResponseWriter, values Values) {
	if values.RequestID != "" {
		w.Header().Set(HeaderRequestID, values.RequestID)
		w.Header().Set(HeaderCorrelationID, values.RequestID)
	}
	if values.TraceID != "" && values.SpanID != "" {
		w.Header().Set(HeaderTraceParent, TraceParent(values))
	}
}

func InjectRequestHeaders(ctx context.Context, req *http.Request) {
	values := FromContext(ctx)
	if values.RequestID != "" {
		req.Header.Set(HeaderRequestID, values.RequestID)
		req.Header.Set(HeaderCorrelationID, values.RequestID)
	}
	if values.TraceID != "" && values.SpanID != "" {
		req.Header.Set(HeaderTraceParent, TraceParent(values))
	}
}

func TraceParent(values Values) string {
	if !isLowerHex(values.TraceID, 32) || !isLowerHex(values.SpanID, 16) {
		return ""
	}
	return "00-" + values.TraceID + "-" + values.SpanID + "-01"
}

func parseTraceParent(header string) (traceID string, spanID string) {
	parts := strings.Split(strings.TrimSpace(header), "-")
	if len(parts) != 4 || parts[0] != "00" {
		return "", ""
	}
	traceID = strings.ToLower(parts[1])
	spanID = strings.ToLower(parts[2])
	if !isLowerHex(traceID, 32) || traceID == "00000000000000000000000000000000" {
		return "", ""
	}
	if !isLowerHex(spanID, 16) || spanID == "0000000000000000" {
		return "", ""
	}
	return traceID, spanID
}

func isLowerHex(value string, length int) bool {
	if len(value) != length {
		return false
	}
	for _, char := range value {
		if (char < '0' || char > '9') && (char < 'a' || char > 'f') {
			return false
		}
	}
	return true
}
