diff --git a/go-backend/internal/http/handler/tunnel_probe_target.go b/go-backend/internal/http/handler/tunnel_probe_target.go index 2110854..d21200f 100644 --- a/go-backend/internal/http/handler/tunnel_probe_target.go +++ b/go-backend/internal/http/handler/tunnel_probe_target.go @@ -4,6 +4,7 @@ import ( "errors" "fmt" "net/netip" + "strconv" "strings" "go-backend/internal/store/model" @@ -135,7 +136,67 @@ func parseTunnelProbeTargetFromRequest(req map[string]interface{}) (tunnelProbeT if req == nil { return defaultTunnelProbeTarget(), false, nil } - return normalizeTunnelProbeTarget(asString(req["probeTargetHost"]), asInt(req["probeTargetPort"], 0)) + rawHost, hasHost := req["probeTargetHost"] + rawPort, hasPort := req["probeTargetPort"] + if !hasHost && !hasPort { + return defaultTunnelProbeTarget(), false, nil + } + host, err := parseTunnelProbeTargetHostValue(rawHost) + if err != nil { + return tunnelProbeTarget{}, false, err + } + port, err := parseTunnelProbeTargetPortValue(rawPort) + if err != nil { + return tunnelProbeTarget{}, false, err + } + return normalizeTunnelProbeTarget(host, port) +} + +func parseTunnelProbeTargetHostValue(raw interface{}) (string, error) { + if raw == nil { + return "", nil + } + host, ok := raw.(string) + if !ok { + return "", errors.New("测试目标 Host 格式无效") + } + if host != strings.TrimSpace(host) { + return "", errors.New("测试目标 Host 不能包含协议或路径") + } + return host, nil +} + +func parseTunnelProbeTargetPortValue(raw interface{}) (int, error) { + if raw == nil { + return 0, nil + } + switch v := raw.(type) { + case float64: + if v != float64(int64(v)) { + return 0, errors.New("测试目标端口必须是整数") + } + return int(v), nil + case string: + if v == "" { + return 0, nil + } + if v != strings.TrimSpace(v) { + return 0, errors.New("测试目标端口必须是整数") + } + port, err := strconv.Atoi(v) + if err != nil { + return 0, errors.New("测试目标端口必须是整数") + } + return port, nil + case int: + return v, nil + case int32: + return int(v), nil + case int64: + return int(v), nil + default: + return 0, errors.New("测试目标端口必须是整数") + } } func effectiveTunnelProbeTarget(tunnel *model.Tunnel) tunnelProbeTarget { diff --git a/go-backend/internal/http/handler/tunnel_probe_target_api_test.go b/go-backend/internal/http/handler/tunnel_probe_target_api_test.go index 01f4e8b..159076f 100644 --- a/go-backend/internal/http/handler/tunnel_probe_target_api_test.go +++ b/go-backend/internal/http/handler/tunnel_probe_target_api_test.go @@ -101,6 +101,53 @@ func TestTunnelUpdateWithoutProbeTargetFieldsPreservesExistingTarget(t *testing. } } +func TestTunnelUpdateRejectsInvalidProbeTargetWithoutClearingExistingTarget(t *testing.T) { + tests := []struct { + name string + probeFields string + }{ + {name: "non numeric port", probeFields: `,"probeTargetPort":"abc"`}, + {name: "fractional port", probeFields: `,"probeTargetPort":443.5`}, + {name: "whitespace host", probeFields: `,"probeTargetHost":" "`}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + h := setupProbeTargetTunnelHandler(t) + seedProbeTargetTunnel(t, h, 80, "existing", "old.example.com", 9443) + body := bytes.NewReader([]byte(`{ + "id":80, + "name":"existing", + "type":1, + "flow":1, + "trafficRatio":1, + "status":1, + "inNodeId":[{"nodeId":10,"protocol":"tls"}] + ` + tt.probeFields + `}`)) + + res := httptest.NewRecorder() + h.tunnelUpdate(res, httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/update", body)) + var payload struct { + Code int `json:"code"` + Msg string `json:"msg"` + } + decodeProbeTargetResponse(t, res, &payload) + if payload.Code == 0 || payload.Msg == "" { + t.Fatalf("expected validation failure, got %+v", payload) + } + + items, err := h.repo.ListTunnels() + if err != nil { + t.Fatalf("list tunnels: %v", err) + } + item := findProbeTargetTunnelItem(t, items, 80) + if item["probeTargetHost"] != "old.example.com" || item["probeTargetPort"] != 9443 { + t.Fatalf("expected invalid probe target to preserve existing target, got %+v", item) + } + }) + } +} + func TestTunnelCreateRejectsInvalidProbeTarget(t *testing.T) { h := setupProbeTargetTunnelHandler(t) body := bytes.NewReader([]byte(`{