93 lines
2.6 KiB
Go
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)(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)
|
|
}
|