package auth

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

func okHandler(w http.ResponseWriter, _ *http.Request) {
	w.WriteHeader(http.StatusOK)
	_, _ = w.Write([]byte("ok"))
}

func TestMiddlewareExempt(t *testing.T) {
	db := openAuthDB(t)
	m := NewMiddleware(NewSessionStore(db, "", 0))
	h := m.Handler(http.HandlerFunc(okHandler))

	for _, p := range []string{
		"/api/system/health",
		"/api/auth/status",
		"/api/auth/webauthn/login/start",
		"/api/auth/webauthn/register/start",
		"/api/auth/recovery",
	} {
		rr := httptest.NewRecorder()
		req := httptest.NewRequest("GET", p, nil)
		h.ServeHTTP(rr, req)
		if rr.Code != http.StatusOK {
			t.Errorf("exempt path %s: got %d, want 200", p, rr.Code)
		}
	}
}

func TestMiddlewareNonAPIPassesThrough(t *testing.T) {
	db := openAuthDB(t)
	m := NewMiddleware(NewSessionStore(db, "", 0))
	h := m.Handler(http.HandlerFunc(okHandler))

	rr := httptest.NewRecorder()
	req := httptest.NewRequest("GET", "/index.html", nil)
	h.ServeHTTP(rr, req)
	if rr.Code != http.StatusOK {
		t.Errorf("non-API path: got %d, want 200", rr.Code)
	}
}

func TestMiddlewareRequiresAuthOnApi(t *testing.T) {
	db := openAuthDB(t)
	m := NewMiddleware(NewSessionStore(db, "", 0))
	h := m.Handler(http.HandlerFunc(okHandler))

	rr := httptest.NewRecorder()
	req := httptest.NewRequest("GET", "/api/trackers/", nil)
	h.ServeHTTP(rr, req)
	if rr.Code != http.StatusUnauthorized {
		t.Errorf("unauthenticated /api/trackers/: got %d, want 401", rr.Code)
	}
}

func TestMiddlewareAcceptsValidSession(t *testing.T) {
	db := openAuthDB(t)
	store := NewSessionStore(db, "", 0)
	m := NewMiddleware(store)

	sess, err := store.Issue(httptest.NewRequest("GET", "/", nil).Context(), "ua", "ip")
	if err != nil {
		t.Fatalf("Issue: %v", err)
	}

	var seenSession *Session
	h := m.Handler(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
		seenSession = SessionFromContext(r.Context())
		w.WriteHeader(http.StatusOK)
	}))

	rr := httptest.NewRecorder()
	req := httptest.NewRequest("GET", "/api/trackers/", nil)
	req.AddCookie(&http.Cookie{Name: "cc_session", Value: sess.Token})
	h.ServeHTTP(rr, req)
	if rr.Code != http.StatusOK {
		t.Fatalf("authenticated /api/trackers/: got %d, want 200", rr.Code)
	}
	if seenSession == nil || seenSession.Token != sess.Token {
		t.Errorf("session not stashed in ctx: %+v", seenSession)
	}
}
