test(httpx): тесты middleware, конверта ошибок и определения IP клиента
This commit is contained in:
@@ -0,0 +1,179 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -140,6 +140,8 @@ func ClientIP(r *http.Request, trustedProxies []*net.IPNet) string {
|
|||||||
return host
|
return host
|
||||||
}
|
}
|
||||||
if xff := r.Header.Get("X-Forwarded-For"); xff != "" {
|
if xff := r.Header.Get("X-Forwarded-For"); xff != "" {
|
||||||
|
// Caddy дописывает в конец цепочки адрес непосредственного клиента,
|
||||||
|
// поэтому при одном доверенном прокси последний элемент и есть клиент.
|
||||||
parts := strings.Split(xff, ",")
|
parts := strings.Split(xff, ",")
|
||||||
candidate := strings.TrimSpace(parts[len(parts)-1])
|
candidate := strings.TrimSpace(parts[len(parts)-1])
|
||||||
if ip := net.ParseIP(candidate); ip != nil {
|
if ip := net.ParseIP(candidate); ip != nil {
|
||||||
|
|||||||
Reference in New Issue
Block a user