package middleware

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

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

	token := "0123456789abcdef0123456789abcdef"
	if err := ValidateCSRFForTests(token, token); err != nil {
		t.Fatalf("expected csrf token to validate: %v", err)
	}
	if err := ValidateCSRFForTests(token, "different789abcdef0123456789abcdef"); err == nil {
		t.Fatal("expected mismatched csrf token to fail")
	}
}

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

	handler := InternalCIDRAllowlist([]string{"10.0.0.0/8"})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
		w.WriteHeader(http.StatusNoContent)
	}))
	request := httptest.NewRequest(http.MethodGet, "/metrics", nil)
	request.RemoteAddr = "10.1.2.3:12345"
	recorder := httptest.NewRecorder()

	handler.ServeHTTP(recorder, request)

	if recorder.Code != http.StatusNoContent {
		t.Fatalf("unexpected status: %d", recorder.Code)
	}
}

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

	handler := InternalCIDRAllowlist([]string{"10.0.0.0/8"})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
		w.WriteHeader(http.StatusNoContent)
	}))
	request := httptest.NewRequest(http.MethodGet, "/metrics", nil)
	request.RemoteAddr = "203.0.113.10:12345"
	recorder := httptest.NewRecorder()

	handler.ServeHTTP(recorder, request)

	if recorder.Code != http.StatusForbidden {
		t.Fatalf("unexpected status: %d", recorder.Code)
	}
}
