package httpx import ( "context" "encoding/json" "log/slog" "net" "net/http" "net/http/httptest" "strings" "testing" ) func TestWriteErrorUsesEnvelope(t *testing.T) { rec := httptest.NewRecorder() WriteError(rec, NewError(CodeNotFound, "no such resource")) if rec.Code != http.StatusNotFound { t.Fatalf("status = %d, want 404", rec.Code) } if got := rec.Header().Get("Content-Type"); got != "application/json; charset=utf-8" { t.Errorf("Content-Type = %q", got) } var body struct { Error Error `json:"error"` } if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil { t.Fatalf("decode envelope: %v", err) } if body.Error.Code != CodeNotFound || body.Error.Message == "" { t.Errorf("unexpected error payload: %+v", body.Error) } } func TestHTTPStatusMapping(t *testing.T) { cases := map[ErrorCode]int{ CodeInternalError: http.StatusInternalServerError, CodeNotFound: http.StatusNotFound, CodeBadRequest: http.StatusBadRequest, CodeRateLimited: http.StatusTooManyRequests, CodeNotReady: http.StatusServiceUnavailable, ErrorCode("unknown.code"): http.StatusInternalServerError, } for code, want := range cases { if got := NewError(code, "").HTTPStatus(); got != want { t.Errorf("status for %q = %d, want %d", code, got, want) } } } func TestRequestIDIsGeneratedAndPropagated(t *testing.T) { var seen string handler := RequestID(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) { seen = RequestIDFrom(r.Context()) })) rec := httptest.NewRecorder() handler.ServeHTTP(rec, httptest.NewRequestWithContext(context.Background(), http.MethodGet, "/", nil)) if seen == "" { t.Fatal("request id is empty inside the handler") } if rec.Header().Get("X-Request-Id") != seen { t.Errorf("header %q does not match context value %q", rec.Header().Get("X-Request-Id"), seen) } } func TestRequestIDIsReusedFromHeader(t *testing.T) { req := httptest.NewRequestWithContext(context.Background(), http.MethodGet, "/", nil) req.Header.Set("X-Request-Id", "client-supplied-id") rec := httptest.NewRecorder() var seen string RequestID(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) { seen = RequestIDFrom(r.Context()) })).ServeHTTP(rec, req) if seen != "client-supplied-id" { t.Errorf("request id = %q, want client-supplied-id", seen) } } func TestRecovererTurnsPanicIntoError(t *testing.T) { logger := slog.New(slog.DiscardHandler) handler := Recoverer(logger)(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { panic("boom") })) rec := httptest.NewRecorder() handler.ServeHTTP(rec, httptest.NewRequestWithContext(context.Background(), http.MethodGet, "/", nil)) if rec.Code != http.StatusInternalServerError { t.Fatalf("status = %d, want 500", rec.Code) } if !json.Valid(rec.Body.Bytes()) { t.Errorf("panic response is not json: %s", rec.Body.String()) } } func TestSecurityHeaders(t *testing.T) { rec := httptest.NewRecorder() SecurityHeaders("files.example.com")(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) })).ServeHTTP(rec, httptest.NewRequestWithContext(context.Background(), http.MethodGet, "/", nil)) headers := rec.Header() for _, name := range []string{"Content-Security-Policy", "Referrer-Policy", "X-Content-Type-Options", "X-Frame-Options", "Cross-Origin-Resource-Policy"} { if headers.Get(name) == "" { t.Errorf("header %s is missing", name) } } if headers.Get("Strict-Transport-Security") != "" { t.Error("HSTS must not be sent over plain http") } } func TestSecurityHeadersAddHSTSBehindProxy(t *testing.T) { req := httptest.NewRequestWithContext(context.Background(), http.MethodGet, "/", nil) req.Header.Set("X-Forwarded-Proto", "https") rec := httptest.NewRecorder() SecurityHeaders("files.example.com")(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) })).ServeHTTP(rec, req) if rec.Header().Get("Strict-Transport-Security") == "" { t.Error("HSTS header is missing for a proxied https request") } } func TestJSONBodyLimitRejectsOversizedBody(t *testing.T) { handler := JSONBodyLimit(16)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { buf := make([]byte, 64) if _, err := r.Body.Read(buf); err != nil { WriteErrorStatus(w, http.StatusRequestEntityTooLarge, CodeBadRequest, "too large") return } w.WriteHeader(http.StatusOK) })) req := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/", strings.NewReader(strings.Repeat("a", 128))) rec := httptest.NewRecorder() handler.ServeHTTP(rec, req) if rec.Code != http.StatusRequestEntityTooLarge { t.Fatalf("status = %d, want 413", rec.Code) } } func TestClientIPPrefersForwardedForFromTrustedProxy(t *testing.T) { _, network, err := net.ParseCIDR("10.0.0.0/8") if err != nil { t.Fatalf("parse cidr: %v", err) } req := httptest.NewRequestWithContext(context.Background(), http.MethodGet, "/", nil) req.RemoteAddr = "10.0.0.5:1234" // Единственный доверенный прокси (Caddy) дописывает клиента в конец цепочки. req.Header.Set("X-Forwarded-For", "203.0.113.9, 198.51.100.7") if got := ClientIP(req, []*net.IPNet{network}); got != "198.51.100.7" { t.Errorf("ClientIP() = %q, want the last hop appended by the trusted proxy", got) } } func TestClientIPIgnoresForwardedForFromUntrustedPeer(t *testing.T) { _, network, err := net.ParseCIDR("10.0.0.0/8") if err != nil { t.Fatalf("parse cidr: %v", err) } req := httptest.NewRequestWithContext(context.Background(), http.MethodGet, "/", nil) req.RemoteAddr = "198.51.100.7:1234" req.Header.Set("X-Forwarded-For", "203.0.113.9") if got := ClientIP(req, []*net.IPNet{network}); got != "198.51.100.7" { t.Errorf("ClientIP() = %q, want the peer address", got) } }