package main
import (
"context"
"errors"
"io"
"log/slog"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
func TestConfiguration(t *testing.T) {
for _, tc := range []struct {
name string
env map[string]string
valid bool
}{
{"missing database", map[string]string{}, false},
{"production needs hosts", map[string]string{"DATABASE_URL": "dbname=test"}, false},
{"production", map[string]string{"DATABASE_URL": "dbname=test", "ALLOWED_HOSTS": "localhost"}, true},
{"compat", map[string]string{"DATABASE_URL": "dbname=test", "APP_PROFILE": "compat"}, true},
{"unknown profile", map[string]string{"DATABASE_URL": "dbname=test", "APP_PROFILE": "typo"}, false},
{"wildcard", map[string]string{"DATABASE_URL": "dbname=test", "ALLOWED_HOSTS": "*"}, false},
{"bad pool", map[string]string{"DATABASE_URL": "dbname=test", "APP_PROFILE": "compat", "DB_MAX_CONNS": "0"}, false},
{"bad admission", map[string]string{"DATABASE_URL": "dbname=test", "APP_PROFILE": "compat", "MAX_INFLIGHT": "99999"}, false},
{"production trace", map[string]string{"DATABASE_URL": "dbname=test", "ALLOWED_HOSTS": "localhost", "SQL_TRACE": "trace.jsonl"}, false},
} {
t.Run(tc.name, func(t *testing.T) {
_, err := loadConfig(func(k string) string { return tc.env[k] })
if (err == nil) != tc.valid {
t.Fatalf("valid=%v err=%v", tc.valid, err)
}
})
}
}
func TestOperationalPolicies(t *testing.T) {
h := &operationalHandler{hosts: map[string]bool{"localhost": true}, slots: make(chan struct{}, 1), logger: slog.New(slog.NewJSONHandler(io.Discard, nil)), ready: func(context.Context) error { return nil }, next: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if _, ok := r.Context().Deadline(); !ok {
t.Error("missing request deadline")
}
w.WriteHeader(204)
})}
request := func(path, host string) *httptest.ResponseRecorder {
r := httptest.NewRequest("GET", "http://localhost"+path, nil)
r.Host = host
r.Header.Set("X-Forwarded-Host", "localhost")
w := httptest.NewRecorder()
h.ServeHTTP(w, r)
return w
}
if w := request("/healthz", "attacker"); w.Code != 400 {
t.Fatal(w.Code)
}
if w := request("/healthz", "localhost:18001"); w.Code != 200 || w.Header().Get("X-Request-ID") == "" {
t.Fatal(w)
}
if w := request("/api/posts", "localhost"); w.Code != 204 {
t.Fatal(w.Code)
}
h.ready = func(context.Context) error { return errors.New("private database error") }
if w := request("/readyz", "localhost"); w.Code != 503 || strings.Contains(w.Body.String(), "private") {
t.Fatal(w)
}
if w := request("/healthz", "localhost"); w.Code != 200 {
t.Fatal(w.Code)
}
h.slots <- struct{}{}
if w := request("/api/posts", "localhost"); w.Code != 503 {
t.Fatal(w.Code)
}
<-h.slots
h.next = http.HandlerFunc(func(http.ResponseWriter, *http.Request) { panic("private panic") })
if w := request("/api/posts", "localhost"); w.Code != 500 || strings.Contains(w.Body.String(), "private") {
t.Fatal(w)
}
}
func TestProductionErrorsAndBodyLimits(t *testing.T) {
a := &Application{production: true}
for _, err := range []error{missingTag{}, duplicateUsers(2), errTitleTooLong} {
w := httptest.NewRecorder()
a.failure(w, err)
if w.Code != 500 || strings.Contains(w.Body.String(), "Traceback") {
t.Fatal(w)
}
}
w := httptest.NewRecorder()
a.failure(w, missing("User"))
if w.Code != 404 {
t.Fatal(w.Code)
}
h := &operationalHandler{next: a, hosts: map[string]bool{"localhost": true}, slots: make(chan struct{}, 1), logger: slog.New(slog.NewJSONHandler(io.Discard, nil))}
for _, length := range []int64{-1, 2 << 20} {
r := httptest.NewRequest("POST", "http://localhost/api/posts", strings.NewReader(strings.Repeat("x", 2<<20)))
r.ContentLength = length
w := httptest.NewRecorder()
h.ServeHTTP(w, r)
if w.Code != 413 {
t.Fatalf("length=%d status=%d", length, w.Code)
}
}
}