package correlation

import (
	"context"
	"net/http"
	"net/http/httptest"
	"testing"
)

func TestFromRequestPreservesRequestAndTraceID(t *testing.T) {
	t.Parallel()

	request := httptest.NewRequest(http.MethodGet, "/", nil)
	request.Header.Set(HeaderRequestID, "req-123")
	request.Header.Set(HeaderTraceParent, "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01")

	values := FromRequest(request)

	if values.RequestID != "req-123" {
		t.Fatalf("RequestID = %q", values.RequestID)
	}
	if values.TraceID != "4bf92f3577b34da6a3ce929d0e0e4736" {
		t.Fatalf("TraceID = %q", values.TraceID)
	}
	if len(values.SpanID) != 16 {
		t.Fatalf("SpanID length = %d", len(values.SpanID))
	}
}

func TestInjectRequestHeaders(t *testing.T) {
	t.Parallel()

	values := Values{
		RequestID: "req-456",
		TraceID:   "4bf92f3577b34da6a3ce929d0e0e4736",
		SpanID:    "00f067aa0ba902b7",
	}
	ctx := WithContext(context.Background(), values)
	request := httptest.NewRequest(http.MethodPost, "https://provider.example.test/payments", nil)

	InjectRequestHeaders(ctx, request)

	if got := request.Header.Get(HeaderRequestID); got != values.RequestID {
		t.Fatalf("request id header = %q", got)
	}
	if got := request.Header.Get(HeaderCorrelationID); got != values.RequestID {
		t.Fatalf("correlation id header = %q", got)
	}
	if got := request.Header.Get(HeaderTraceParent); got != "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01" {
		t.Fatalf("traceparent header = %q", got)
	}
}

func TestInvalidTraceParentGeneratesNewTrace(t *testing.T) {
	t.Parallel()

	request := httptest.NewRequest(http.MethodGet, "/", nil)
	request.Header.Set(HeaderTraceParent, "invalid")

	values := FromRequest(request)

	if len(values.TraceID) != 32 {
		t.Fatalf("expected generated trace id, got %q", values.TraceID)
	}
	if TraceParent(values) == "" {
		t.Fatal("expected generated values to produce traceparent")
	}
}
