fix: 校验端默认延迟不超过调度上限
This commit is contained in:
@@ -233,7 +233,7 @@ func (h *Handler) handleEndpointCreate(w http.ResponseWriter, r *http.Request) {
|
||||
if req.DefaultDelaySeconds != nil {
|
||||
delaySec = *req.DefaultDelaySeconds
|
||||
}
|
||||
if errMsg := validateEndpointFields(req.ID, req.Name, req.Remark, req.LoginPassword, req.TalkPassword, delaySec); errMsg != "" {
|
||||
if errMsg := validateEndpointFields(req.ID, req.Name, req.Remark, req.LoginPassword, req.TalkPassword, delaySec, h.maxScheduleSeconds()); errMsg != "" {
|
||||
h.audit(actorString(p), "endpoint_create", req.ID, "bad_request", ip)
|
||||
httpx.WriteError(w, http.StatusBadRequest, "bad_request", errMsg)
|
||||
return
|
||||
@@ -341,10 +341,12 @@ func (h *Handler) handleEndpointPatch(w http.ResponseWriter, r *http.Request) {
|
||||
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "备注过长")
|
||||
return
|
||||
}
|
||||
if req.DefaultDelaySeconds != nil && *req.DefaultDelaySeconds < 0 {
|
||||
h.audit(actorString(p), "endpoint_patch", id, "bad_request", ip)
|
||||
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "默认延迟无效")
|
||||
return
|
||||
if req.DefaultDelaySeconds != nil {
|
||||
if msg := validateDelaySeconds(*req.DefaultDelaySeconds, h.maxScheduleSeconds()); msg != "" {
|
||||
h.audit(actorString(p), "endpoint_patch", id, "bad_request", ip)
|
||||
httpx.WriteError(w, http.StatusBadRequest, "bad_request", msg)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
hasMeta := req.Name != nil || req.Remark != nil || req.DefaultDelaySeconds != nil
|
||||
@@ -614,7 +616,7 @@ func (h *Handler) handleEndpointUnlock(w http.ResponseWriter, r *http.Request) {
|
||||
httpx.WriteOK(w, map[string]any{})
|
||||
}
|
||||
|
||||
func validateEndpointFields(id, name, remark, loginPW, talkPW string, delaySec int64) string {
|
||||
func validateEndpointFields(id, name, remark, loginPW, talkPW string, delaySec, maxDelaySec int64) string {
|
||||
if id != "" && !protocol.ValidEndpointID(id) {
|
||||
return "编号不合法"
|
||||
}
|
||||
@@ -633,12 +635,26 @@ func validateEndpointFields(id, name, remark, loginPW, talkPW string, delaySec i
|
||||
if !protocol.ValidTalkPassword(talkPW) {
|
||||
return "对话密码不合法"
|
||||
}
|
||||
if msg := validateDelaySeconds(delaySec, maxDelaySec); msg != "" {
|
||||
return msg
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func validateDelaySeconds(delaySec, maxDelaySec int64) string {
|
||||
if delaySec < 0 {
|
||||
return "默认延迟无效"
|
||||
}
|
||||
if maxDelaySec > 0 && delaySec > maxDelaySec {
|
||||
return "默认延迟无效"
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *Handler) maxScheduleSeconds() int64 {
|
||||
return int64(h.cfg.Limits.MaxScheduleSeconds)
|
||||
}
|
||||
|
||||
func writeCSVValidationError(w http.ResponseWriter, errs []csvLineError) {
|
||||
httpx.WriteJSON(w, http.StatusBadRequest, httpx.Envelope{
|
||||
OK: false,
|
||||
|
||||
@@ -244,7 +244,7 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr
|
||||
}
|
||||
delaySec = n
|
||||
}
|
||||
if msg := validateEndpointFields(id, name, remark, loginPW, talkPW, delaySec); msg != "" {
|
||||
if msg := validateEndpointFields(id, name, remark, loginPW, talkPW, delaySec, h.maxScheduleSeconds()); msg != "" {
|
||||
errs = append(errs, csvLineError{Line: line, Reason: msg})
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -0,0 +1,61 @@
|
||||
package admin_test
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestCreateRejectsDefaultDelayOverMax(t *testing.T) {
|
||||
_, srv, client, _, _ := setupEndpoints(t)
|
||||
body := `{"id":"dly-c","name":"延迟","login_password":"password1","default_delay_seconds":31536001}`
|
||||
res := postJSON(t, client, srv.URL+"/api/admin/endpoints", body, csrfHeaders())
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != http.StatusBadRequest {
|
||||
t.Fatalf("want 400 got %d %+v", res.StatusCode, env)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPatchRejectsDefaultDelayOverMax(t *testing.T) {
|
||||
_, srv, client, _, _ := setupEndpoints(t)
|
||||
res := postJSON(t, client, srv.URL+"/api/admin/endpoints",
|
||||
`{"id":"dly-p","name":"延迟","login_password":"password1"}`, csrfHeaders())
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("create: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
res = doReq(t, client, http.MethodPatch, srv.URL+"/api/admin/endpoints/dly-p",
|
||||
`{"default_delay_seconds":31536001}`, csrfHeaders())
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != http.StatusBadRequest {
|
||||
t.Fatalf("want 400 got %d %+v", res.StatusCode, env)
|
||||
}
|
||||
}
|
||||
|
||||
func TestImportRejectsDefaultDelayOverMax(t *testing.T) {
|
||||
_, srv, client, _, _ := setupEndpoints(t)
|
||||
csvBody := "" +
|
||||
"id,name,login_password,talk_password,default_delay_seconds,remark\n" +
|
||||
"dly-i,甲,password1,,31536001,\n"
|
||||
res := postImportCSV(t, client, srv.URL, csvBody)
|
||||
raw, _ := io.ReadAll(res.Body)
|
||||
_ = res.Body.Close()
|
||||
if res.StatusCode != http.StatusBadRequest {
|
||||
t.Fatalf("want 400 got %d body=%s", res.StatusCode, raw)
|
||||
}
|
||||
var parsed struct {
|
||||
Data *struct {
|
||||
Errors []struct {
|
||||
Line int `json:"line"`
|
||||
Reason string `json:"reason"`
|
||||
} `json:"errors"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal(raw, &parsed); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if parsed.Data == nil || len(parsed.Data.Errors) == 0 {
|
||||
t.Fatalf("want row error, body=%s", raw)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user