fix: 限制管理接口请求体大小与读取时间

This commit is contained in:
Nixevol
2026-09-30 16:22:48 +08:00
parent b40ef5c548
commit 1a4bf6f185
7 changed files with 170 additions and 8 deletions
+14 -3
View File
@@ -37,6 +37,11 @@ func (h *Handler) handleEndpointImport(w http.ResponseWriter, r *http.Request) {
raw, err := readImportCSV(r)
if err != nil {
if httpx.IsBodyTooLarge(err) {
h.audit(actorString(p), "endpoint_import", "", "payload_too_large", ip)
httpx.WriteError(w, http.StatusRequestEntityTooLarge, "payload_too_large", "请求体过大")
return
}
h.audit(actorString(p), "endpoint_import", "", "bad_request", ip)
httpx.WriteError(w, http.StatusBadRequest, "bad_request", err.Error())
return
@@ -91,9 +96,12 @@ func readImportCSV(r *http.Request) ([]byte, error) {
}
name := part.FormName()
if name == "file" || name == "" {
b, readErr := io.ReadAll(io.LimitReader(part, 8<<20))
b, readErr := io.ReadAll(part)
_ = part.Close()
if readErr != nil {
if httpx.IsBodyTooLarge(readErr) {
return nil, readErr
}
return nil, errBadRequest("读取文件失败")
}
return b, nil
@@ -102,9 +110,12 @@ func readImportCSV(r *http.Request) ([]byte, error) {
}
return nil, errBadRequest("缺少 file 字段")
default:
// text/csv 或未标明时按原始体
b, readErr := io.ReadAll(io.LimitReader(r.Body, 8<<20))
// text/csv 或未标明时按原始体;大小由 ServeHTTP 的 MaxBytesReader 限制。
b, readErr := io.ReadAll(r.Body)
if readErr != nil {
if httpx.IsBodyTooLarge(readErr) {
return nil, readErr
}
return nil, errBadRequest("读取 CSV 失败")
}
return b, nil
+53
View File
@@ -0,0 +1,53 @@
package admin_test
import (
"bytes"
"io"
"net/http"
"strings"
"testing"
)
func TestLoginRejectsOversizeBody(t *testing.T) {
_, srv, client, _ := setup(t)
body := `{"username":"admin","password":"` + strings.Repeat("a", 1<<20) + `"}`
req, err := http.NewRequest(http.MethodPost, srv.URL+"/api/admin/login", strings.NewReader(body))
if err != nil {
t.Fatal(err)
}
req.Header.Set("Content-Type", "application/json")
res, err := client.Do(req)
if err != nil {
t.Fatal(err)
}
env := decodeEnv(t, res)
if res.StatusCode != http.StatusRequestEntityTooLarge {
t.Fatalf("want 413 got %d env=%+v", res.StatusCode, env)
}
if env.Error == nil || env.Error.Code != "payload_too_large" {
t.Fatalf("want payload_too_large got %+v", env.Error)
}
}
func TestImportRejectsOversizeCSV(t *testing.T) {
_, srv, client, _, _ := setupEndpoints(t)
payload := bytes.Repeat([]byte("x"), 8<<20+1)
req, err := http.NewRequest(http.MethodPost, srv.URL+"/api/admin/endpoints/import", bytes.NewReader(payload))
if err != nil {
t.Fatal(err)
}
req.Header.Set("Content-Type", "text/csv")
req.Header.Set("X-Nixmsg-Request", "1")
res, err := client.Do(req)
if err != nil {
t.Fatal(err)
}
raw, _ := io.ReadAll(res.Body)
_ = res.Body.Close()
if res.StatusCode != http.StatusRequestEntityTooLarge {
t.Fatalf("want 413 got %d body=%s", res.StatusCode, raw)
}
if !strings.Contains(string(raw), "payload_too_large") {
t.Fatalf("want payload_too_large in body: %s", raw)
}
}
+35 -1
View File
@@ -1,6 +1,7 @@
package admin
import (
"errors"
"log/slog"
"net"
"net/http"
@@ -11,6 +12,7 @@ import (
"git.asio.asia/nixevol/NixMsg/internal/app/identity"
"git.asio.asia/nixevol/NixMsg/internal/auth"
"git.asio.asia/nixevol/NixMsg/internal/config"
"git.asio.asia/nixevol/NixMsg/internal/httpx"
"git.asio.asia/nixevol/NixMsg/internal/store"
)
@@ -23,6 +25,13 @@ const (
defaultSessionTTL = 12 * time.Hour
minPasswordLen = 12
lastUsedMinGap = time.Minute
maxLoginBodyBytes = 8 << 10
maxJSONBodyBytes = 1 << 20
maxImportBodyBytes = 8 << 20
defaultReadFor = 15 * time.Second
loginReadFor = 10 * time.Second
importReadFor = 2 * time.Minute
)
// Deps 是管理 Handler 的依赖。
@@ -129,11 +138,36 @@ func New(d Deps) *Handler {
return h
}
// ServeHTTP 实现 http.Handler。
// ServeHTTP 实现 http.Handler。按路由限制请求体大小与读截止时间。
func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
limit, readFor := requestBodyBudget(r)
if err := http.NewResponseController(w).SetReadDeadline(time.Now().Add(readFor)); err != nil && !errors.Is(err, http.ErrNotSupported) {
// 测试用 ResponseRecorder 或不支持截止时间的封装:忽略。
}
if r.Body != nil {
r.Body = http.MaxBytesReader(w, r.Body, limit)
}
h.mux.ServeHTTP(w, r)
}
func requestBodyBudget(r *http.Request) (int64, time.Duration) {
if r.Method == http.MethodPost && r.URL.Path == "/api/admin/endpoints/import" {
return maxImportBodyBytes, importReadFor
}
if r.Method == http.MethodPost && r.URL.Path == "/api/admin/login" {
return maxLoginBodyBytes, loginReadFor
}
return maxJSONBodyBytes, defaultReadFor
}
func writeDecodeError(w http.ResponseWriter, err error) {
if httpx.IsBodyTooLarge(err) {
httpx.WriteError(w, http.StatusRequestEntityTooLarge, "payload_too_large", "请求体过大")
return
}
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效")
}
func (h *Handler) routes() {
// 公开
h.mux.HandleFunc("POST /api/admin/login", h.handleLogin)
+2 -2
View File
@@ -35,7 +35,7 @@ func (h *Handler) handleLogin(w http.ResponseWriter, r *http.Request) {
Password string `json:"password"`
}
if err := httpx.DecodeJSON(r, &req); err != nil {
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效")
writeDecodeError(w, err)
return
}
if req.Username != adminUsername {
@@ -117,7 +117,7 @@ func (h *Handler) handlePassword(w http.ResponseWriter, r *http.Request) {
}
if err := httpx.DecodeJSON(r, &req); err != nil {
h.audit(actorString(p), "password_change", "", "bad_request", ip)
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效")
writeDecodeError(w, err)
return
}
if len(req.NewPassword) < minPasswordLen {