Files
cmroubao/backend-api/internal/transport/httpapi/router_test.go
T

288 lines
7.8 KiB
Go
Raw Normal View History

package httpapi
import (
"context"
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"regexp"
"strings"
"testing"
"time"
"cmroubao/backend-api/internal/domain"
"github.com/gin-gonic/gin"
)
type fakePinger struct {
err error
}
func (p fakePinger) PingContext(context.Context) error {
return p.err
}
func TestHealthzReturnsStableHealthyResponse(t *testing.T) {
router, err := newTestRouter(fakePinger{}, discardEvent)
if err != nil {
t.Fatalf("NewRouter() error = %v", err)
}
response := performRequest(t, router, http.MethodGet, "/healthz")
if response.Code != http.StatusOK {
t.Fatalf("status = %d", response.Code)
}
assertJSON(t, response, map[string]any{"status": "ok"})
if response.Header().Get("Cache-Control") != "no-store" {
t.Fatalf("Cache-Control = %q", response.Header().Get("Cache-Control"))
}
if response.Header().Get("Access-Control-Allow-Origin") != "" {
t.Fatal("default CORS must remain disabled")
}
}
func TestHealthzReturns503WithoutLeakingDatabaseError(t *testing.T) {
router, err := newTestRouter(fakePinger{
err: errors.New("private database path and driver details"),
}, discardEvent)
if err != nil {
t.Fatalf("NewRouter() error = %v", err)
}
response := performRequest(t, router, http.MethodGet, "/healthz")
if response.Code != http.StatusServiceUnavailable {
t.Fatalf("status = %d", response.Code)
}
if strings.Contains(response.Body.String(), "private database") {
t.Fatal("health response leaked the database error")
}
assertJSON(t, response, map[string]any{"status": "unavailable"})
}
func TestUnknownRouteAndMethodUseStableErrors(t *testing.T) {
router, err := newTestRouter(fakePinger{}, discardEvent)
if err != nil {
t.Fatalf("NewRouter() error = %v", err)
}
notFound := performRequest(t, router, http.MethodGet, "/missing")
if notFound.Code != http.StatusNotFound {
t.Fatalf("not found status = %d", notFound.Code)
}
assertErrorCode(t, notFound, "NOT_FOUND")
notAllowed := performRequest(t, router, http.MethodPost, "/healthz")
if notAllowed.Code != http.StatusMethodNotAllowed {
t.Fatalf("not allowed status = %d", notAllowed.Code)
}
assertErrorCode(t, notAllowed, "METHOD_NOT_ALLOWED")
}
func TestSafeRecoveryReturnsStableErrorWithoutLoggingRequestHeaders(t *testing.T) {
var events []string
router, err := newTestRouter(
panicPinger{},
func(event string) {
events = append(events, event)
},
)
if err != nil {
t.Fatalf("NewRouter() error = %v", err)
}
request := httptest.NewRequest(http.MethodGet, "/healthz", nil)
request.Header.Set("Authorization", "Bearer private-token")
request.Header.Set("Cookie", "session=private-cookie")
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code != http.StatusInternalServerError {
t.Fatalf("status = %d", response.Code)
}
assertErrorCode(t, response, "INTERNAL_ERROR")
if len(events) != 1 || events[0] != "HTTP handler panic recovered" {
t.Fatalf("events = %#v", events)
}
eventText := strings.Join(events, " ")
if strings.Contains(eventText, "private") ||
strings.Contains(eventText, "Bearer") ||
strings.Contains(eventText, "session") {
t.Fatalf("event log leaked request or panic content")
}
}
func TestRouterRequiresDependencies(t *testing.T) {
valid := RouterDependencies{
Database: fakePinger{},
RegisterPublicRoutes: discardRoutes,
RegisterAdminRoutes: discardRoutes,
RegisterDeviceRoutes: discardRoutes,
AdminSessions: allowAdminAuthenticator{},
DeviceAccess: allowAdminAuthenticator{},
LogEvent: discardEvent,
}
missingDatabase := valid
missingDatabase.Database = nil
if _, err := NewRouter(missingDatabase); err == nil {
t.Fatal("NewRouter(nil database) error = nil")
}
missingRoutes := valid
missingRoutes.RegisterAdminRoutes = nil
if _, err := NewRouter(missingRoutes); err == nil {
t.Fatal("NewRouter(nil routes) error = nil")
}
missingPublicRoutes := valid
missingPublicRoutes.RegisterPublicRoutes = nil
if _, err := NewRouter(missingPublicRoutes); err == nil {
t.Fatal("NewRouter(nil public routes) error = nil")
}
missingDeviceRoutes := valid
missingDeviceRoutes.RegisterDeviceRoutes = nil
if _, err := NewRouter(missingDeviceRoutes); err == nil {
t.Fatal("NewRouter(nil device routes) error = nil")
}
missingAuth := valid
missingAuth.AdminSessions = nil
if _, err := NewRouter(missingAuth); err == nil {
t.Fatal("NewRouter(nil admin auth) error = nil")
}
missingDeviceAuth := valid
missingDeviceAuth.DeviceAccess = nil
if _, err := NewRouter(missingDeviceAuth); err == nil {
t.Fatal("NewRouter(nil device auth) error = nil")
}
missingLogger := valid
missingLogger.LogEvent = nil
if _, err := NewRouter(missingLogger); err == nil {
t.Fatal("NewRouter(nil logger) error = nil")
}
}
func newTestRouter(
database DatabasePinger,
logEvent EventLogger,
) (http.Handler, error) {
return NewRouter(RouterDependencies{
Database: database,
RegisterPublicRoutes: discardRoutes,
RegisterAdminRoutes: discardRoutes,
RegisterDeviceRoutes: discardRoutes,
AdminSessions: allowAdminAuthenticator{},
DeviceAccess: allowAdminAuthenticator{},
LogEvent: logEvent,
})
}
func performRequest(
t *testing.T,
handler http.Handler,
method string,
path string,
) *httptest.ResponseRecorder {
t.Helper()
request := httptest.NewRequest(method, path, nil)
response := httptest.NewRecorder()
handler.ServeHTTP(response, request)
return response
}
func assertJSON(
t *testing.T,
response *httptest.ResponseRecorder,
want map[string]any,
) {
t.Helper()
var got map[string]any
if err := json.Unmarshal(response.Body.Bytes(), &got); err != nil {
t.Fatalf("decode JSON: %v", err)
}
if len(got) != len(want) {
t.Fatalf("JSON = %#v, want %#v", got, want)
}
for key, wantValue := range want {
if got[key] != wantValue {
t.Fatalf("JSON[%q] = %#v, want %#v", key, got[key], wantValue)
}
}
}
func assertErrorCode(
t *testing.T,
response *httptest.ResponseRecorder,
want string,
) {
t.Helper()
var body struct {
Error struct {
Code string `json:"code"`
Retryable bool `json:"retryable"`
Details map[string]any `json:"details"`
} `json:"error"`
RequestID string `json:"request_id"`
}
if err := json.Unmarshal(response.Body.Bytes(), &body); err != nil {
t.Fatalf("decode error JSON: %v", err)
}
if body.Error.Code != want {
t.Fatalf("error code = %q, want %q", body.Error.Code, want)
}
if body.Error.Retryable {
t.Fatal("retryable = true")
}
if body.Error.Details == nil || len(body.Error.Details) != 0 {
t.Fatalf("details = %#v", body.Error.Details)
}
if !requestIDPattern.MatchString(body.RequestID) {
t.Fatalf("request_id = %q", body.RequestID)
}
if response.Header().Get(requestIDHeader) != body.RequestID {
t.Fatalf("request ID header does not match body")
}
}
type panicPinger struct{}
func (panicPinger) PingContext(context.Context) error {
panic("private failure detail")
}
func discardEvent(string) {}
func discardRoutes(gin.IRoutes) error { return nil }
type allowAdminAuthenticator struct{}
func (allowAdminAuthenticator) AuthenticateAdmin(
context.Context,
string,
) (domain.AuthPrincipal, error) {
return domain.AuthPrincipal{
UserID: "00000000-0000-4000-8000-000000000099",
Username: "admin",
Role: domain.UserRoleAdmin,
SessionID: "admin-session",
ExpiresAt: time.Now().Add(time.Hour),
}, nil
}
func (allowAdminAuthenticator) AuthenticateAccessToken(
context.Context,
string,
) (domain.AuthPrincipal, error) {
return domain.AuthPrincipal{
UserID: "00000000-0000-4000-8000-000000000098",
Username: "buyer",
Role: domain.UserRoleBuyer,
DeviceID: "00000000-0000-4000-8000-000000000097",
ExpiresAt: time.Now().Add(time.Hour),
}, nil
}
var requestIDPattern = regexp.MustCompile(
`^[0-9a-f]{8}-[0-9a-f]{4}-4[0-9a-f]{3}-[89ab][0-9a-f]{3}-[0-9a-f]{12}$`,
)