feat(api): auth, request_id, and request-logging middleware
This commit is contained in:
@@ -0,0 +1,92 @@
|
||||
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)
|
||||
}
|
||||
Reference in New Issue
Block a user