package api_test import ( "io" "log/slog" "net/http" "net/http/httptest" "testing" "github.com/shcizo/package-updater/internal/api" "github.com/shcizo/package-updater/internal/logging" "github.com/stretchr/testify/require" ) func newTestLogger() *slog.Logger { return slog.New(slog.NewJSONHandler(io.Discard, nil)) } func TestAuth_AllowsMatchingToken(t *testing.T) { called := false h := api.Auth("secret")(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { called = true w.WriteHeader(http.StatusOK) })) req := httptest.NewRequest(http.MethodPost, "/update", nil) req.Header.Set("Authorization", "Bearer secret") w := httptest.NewRecorder() h.ServeHTTP(w, req) require.True(t, called) require.Equal(t, http.StatusOK, w.Code) } func TestAuth_Rejects(t *testing.T) { cases := []struct{ name, header string }{ {"missing", ""}, {"wrong scheme", "Token secret"}, {"wrong value", "Bearer nope"}, {"empty bearer", "Bearer "}, } for _, c := range cases { t.Run(c.name, func(t *testing.T) { h := api.Auth("secret")(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { t.Fatal("handler must not be called") })) req := httptest.NewRequest(http.MethodPost, "/update", nil) if c.header != "" { req.Header.Set("Authorization", c.header) } w := httptest.NewRecorder() h.ServeHTTP(w, req) require.Equal(t, http.StatusUnauthorized, w.Code) }) } } func TestRequestID_GeneratesIfMissing(t *testing.T) { var seenID string h := api.RequestID(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) { seenID = logging.RequestIDFrom(r.Context()) })) req := httptest.NewRequest(http.MethodPost, "/update", nil) w := httptest.NewRecorder() h.ServeHTTP(w, req) require.NotEmpty(t, seenID) require.Equal(t, seenID, w.Header().Get("X-Request-ID")) } func TestRequestID_UsesIncoming(t *testing.T) { var seenID string h := api.RequestID(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) { seenID = logging.RequestIDFrom(r.Context()) })) req := httptest.NewRequest(http.MethodPost, "/update", nil) req.Header.Set("X-Request-ID", "given-id") w := httptest.NewRecorder() h.ServeHTTP(w, req) require.Equal(t, "given-id", seenID) } func TestRequestLogger_LogsAndDelegates(t *testing.T) { logger := newTestLogger() called := false h := api.RequestLogger(logger)(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { called = true w.WriteHeader(http.StatusTeapot) })) req := httptest.NewRequest(http.MethodPost, "/update", nil) w := httptest.NewRecorder() h.ServeHTTP(w, req) require.True(t, called) require.Equal(t, http.StatusTeapot, w.Code) }