fix(forward): split proxy protocol directions

Closes #520
This commit is contained in:
sagit
2026-06-21 21:03:47 +08:00
committed by GitHub
parent 3ce320da5a
commit 82f6047506
13 changed files with 505 additions and 259 deletions
@@ -1757,6 +1757,7 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
services := make([]map[string]interface{}, 0, 2)
targets := splitRemoteTargets(forward.RemoteAddr)
strategy := strings.TrimSpace(forward.Strategy)
proxyProtocolReceive, proxyProtocolSend := normalizeForwardProxyProtocol(forward.ProxyProtocol, forward.ProxyProtocolReceive, forward.ProxyProtocolSend)
if strategy == "" {
strategy = "fifo"
}
@@ -1801,12 +1802,16 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
if runtimeLimiters.TrafficLimiter != "" {
service["limiter"] = runtimeLimiters.TrafficLimiter
}
if forward.ProxyProtocol > 0 {
if proxyProtocolReceive > 0 {
serviceMetadata := ensureServiceMetadata(service)
serviceMetadata["proxyProtocol"] = proxyProtocolReceive
}
if proxyProtocolSend > 0 {
handlerConfig := service["handler"].(map[string]interface{})
if handlerConfig["metadata"] == nil {
handlerConfig["metadata"] = map[string]interface{}{}
}
handlerConfig["metadata"].(map[string]interface{})["proxyProtocol"] = forward.ProxyProtocol
handlerConfig["metadata"].(map[string]interface{})["proxyProtocol"] = proxyProtocolSend
}
if protocol == "udp" {
listenerMetadata := map[string]interface{}{
@@ -1819,10 +1824,8 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
service["handler"].(map[string]interface{})["chain"] = fmt.Sprintf("chains_%d", forward.TunnelID)
}
if tunnel != nil && tunnel.Type == 1 && strings.TrimSpace(node.InterfaceName) != "" {
if service["metadata"] == nil {
service["metadata"] = map[string]interface{}{}
}
service["metadata"].(map[string]interface{})["interface"] = node.InterfaceName
serviceMetadata := ensureServiceMetadata(service)
serviceMetadata["interface"] = node.InterfaceName
}
services = append(services, service)
}
@@ -1841,6 +1844,25 @@ func buildForwarderNodes(targets []string) []map[string]interface{} {
return nodes
}
func ensureServiceMetadata(service map[string]interface{}) map[string]interface{} {
if service["metadata"] == nil {
service["metadata"] = map[string]interface{}{}
}
metadata, ok := service["metadata"].(map[string]interface{})
if !ok {
metadata = map[string]interface{}{}
service["metadata"] = metadata
}
return metadata
}
func normalizeForwardProxyProtocol(legacy, receive, send int) (int, int) {
if send == 0 && legacy > 0 {
send = legacy
}
return receive, send
}
func processServerAddress(serverAddr string) string {
serverAddr = normalizeServerAddressInput(serverAddr)
if serverAddr == "" {
@@ -9,14 +9,15 @@ import (
"go-backend/internal/store/repo"
)
func TestBuildForwardServiceConfigsSendsProxyProtocolToForwardHandler(t *testing.T) {
func TestBuildForwardServiceConfigsAppliesProxyProtocolReceiveAndSendIndependently(t *testing.T) {
forward := &forwardRecord{
ID: 1,
UserID: 2,
TunnelID: 3,
RemoteAddr: "1.1.1.1:443",
Strategy: "fifo",
ProxyProtocol: 2,
ID: 1,
UserID: 2,
TunnelID: 3,
RemoteAddr: "1.1.1.1:443",
Strategy: "fifo",
ProxyProtocolReceive: 1,
ProxyProtocolSend: 2,
}
tunnel := &tunnelRecord{Type: 1}
node := &nodeRecord{
@@ -38,8 +39,8 @@ func TestBuildForwardServiceConfigsSendsProxyProtocolToForwardHandler(t *testing
if serviceMetadata["interface"] != "eth0" {
t.Fatalf("expected interface metadata eth0, got %v", serviceMetadata["interface"])
}
if _, ok := serviceMetadata["proxyProtocol"]; ok {
t.Fatalf("proxyProtocol should not be listener metadata: %v", serviceMetadata)
if serviceMetadata["proxyProtocol"] != 1 {
t.Fatalf("expected service proxyProtocol 1 for receive mode, got %v", serviceMetadata["proxyProtocol"])
}
handlerConfig, ok := service["handler"].(map[string]interface{})
@@ -56,6 +57,46 @@ func TestBuildForwardServiceConfigsSendsProxyProtocolToForwardHandler(t *testing
}
}
func TestBuildForwardServiceConfigsKeepsLegacyProxyProtocolAsSend(t *testing.T) {
forward := &forwardRecord{
ID: 1,
UserID: 2,
TunnelID: 3,
RemoteAddr: "1.1.1.1:443",
Strategy: "fifo",
ProxyProtocol: 2,
}
tunnel := &tunnelRecord{Type: 1}
node := &nodeRecord{
TCPListenAddr: "0.0.0.0",
UDPListenAddr: "0.0.0.0",
}
services := buildForwardServiceConfigs("1_2_3", forward, tunnel, node, 4001, "", forwardRuntimeLimiters{})
if len(services) != 2 {
t.Fatalf("expected 2 services, got %d", len(services))
}
for _, service := range services {
serviceMetadata, _ := service["metadata"].(map[string]interface{})
if _, ok := serviceMetadata["proxyProtocol"]; ok {
t.Fatalf("legacy proxyProtocol should not enable receive mode: %v", serviceMetadata)
}
handlerConfig, ok := service["handler"].(map[string]interface{})
if !ok {
t.Fatalf("expected handler config map, got %T", service["handler"])
}
handlerMetadata, ok := handlerConfig["metadata"].(map[string]interface{})
if !ok {
t.Fatalf("expected handler metadata map, got %T", handlerConfig["metadata"])
}
if handlerMetadata["proxyProtocol"] != 2 {
t.Fatalf("expected legacy proxyProtocol to send version 2, got %v", handlerMetadata["proxyProtocol"])
}
}
}
func TestRollbackForwardMutationRestoresProxyProtocol(t *testing.T) {
r, err := repo.Open(":memory:")
if err != nil {
@@ -2143,8 +2143,10 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
ipMaxConn = 0
}
proxyProtocol := asInt(req["proxyProtocol"], 0)
proxyProtocolReceive := asInt(req["proxyProtocolReceive"], 0)
proxyProtocolSend := asInt(req["proxyProtocolSend"], proxyProtocol)
forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port, inIp, nullableInt(speedID), maxConn, ipMaxConn, nullableInt(ipSpeedID), proxyProtocol)
forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port, inIp, nullableInt(speedID), maxConn, ipMaxConn, nullableInt(ipSpeedID), proxyProtocol, proxyProtocolReceive, proxyProtocolSend)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
@@ -2333,8 +2335,10 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
ipMaxConn = 0
}
proxyProtocol := asInt(req["proxyProtocol"], forward.ProxyProtocol)
proxyProtocolReceive := asInt(req["proxyProtocolReceive"], forward.ProxyProtocolReceive)
proxyProtocolSend := asInt(req["proxyProtocolSend"], forward.ProxyProtocolSend)
if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now, newSpeedID, maxConn, ipMaxConn, newIPSpeedID, proxyProtocol); err != nil {
if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now, newSpeedID, maxConn, ipMaxConn, newIPSpeedID, proxyProtocol, proxyProtocolReceive, proxyProtocolSend); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
@@ -4722,7 +4726,7 @@ func (h *Handler) rollbackForwardMutation(oldForward *forwardRecord, oldPorts []
h.repo.RollbackForwardFields(
oldForward.ID, oldForward.UserID, oldForward.UserName, oldForward.Name,
oldForward.TunnelID, oldForward.RemoteAddr, oldForward.Strategy, oldForward.Status,
oldForward.SpeedID, oldForward.MaxConn, oldForward.IPMaxConn, oldForward.IPSpeedID, oldForward.ProxyProtocol,
oldForward.SpeedID, oldForward.MaxConn, oldForward.IPMaxConn, oldForward.IPSpeedID, oldForward.ProxyProtocol, oldForward.ProxyProtocolReceive, oldForward.ProxyProtocolSend,
time.Now().UnixMilli(),
)
@@ -490,7 +490,7 @@ func seedForwardForNftables(t *testing.T, h *Handler, tunnelID, nodeID int64, re
now := time.Now().UnixMilli()
forwardID, err := h.repo.CreateForwardTx(
1, "admin", "nft-forward", tunnelID, remoteAddr, "fifo", now, 1,
[]int64{nodeID}, 20000, "", nil, 0, 0, nil, 0,
[]int64{nodeID}, 20000, "", nil, 0, 0, nil, 0, 0, 0,
)
if err != nil {
t.Fatalf("create forward: %v", err)