mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-07 18:26:37 +08:00
fix: reject malformed probe target updates
This commit is contained in:
@@ -4,6 +4,7 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"go-backend/internal/store/model"
|
"go-backend/internal/store/model"
|
||||||
@@ -135,7 +136,67 @@ func parseTunnelProbeTargetFromRequest(req map[string]interface{}) (tunnelProbeT
|
|||||||
if req == nil {
|
if req == nil {
|
||||||
return defaultTunnelProbeTarget(), false, 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 {
|
func effectiveTunnelProbeTarget(tunnel *model.Tunnel) tunnelProbeTarget {
|
||||||
|
|||||||
@@ -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) {
|
func TestTunnelCreateRejectsInvalidProbeTarget(t *testing.T) {
|
||||||
h := setupProbeTargetTunnelHandler(t)
|
h := setupProbeTargetTunnelHandler(t)
|
||||||
body := bytes.NewReader([]byte(`{
|
body := bytes.NewReader([]byte(`{
|
||||||
|
|||||||
Reference in New Issue
Block a user