Files

93 lines
2.6 KiB
Go

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, nil)(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)
}