fix(backend): normalize strategy data and proxy ip parsing

This commit is contained in:
sagit
2026-02-13 09:42:38 +00:00
parent f01c0481cd
commit cf6294a77d
7 changed files with 71 additions and 10 deletions
@@ -114,7 +114,7 @@ func (h *Handler) ensureTunnelPermission(userID int64, roleID int, tunnelID int6
func (h *Handler) getForwardRecord(forwardID int64) (*forwardRecord, error) { func (h *Handler) getForwardRecord(forwardID int64) (*forwardRecord, error) {
row := h.repo.DB().QueryRow(` row := h.repo.DB().QueryRow(`
SELECT id, user_id, user_name, name, tunnel_id, remote_addr, strategy, status SELECT id, user_id, user_name, name, tunnel_id, remote_addr, COALESCE(strategy, 'fifo'), status
FROM forward WHERE id = ? LIMIT 1 FROM forward WHERE id = ? LIMIT 1
`, forwardID) `, forwardID)
var fr forwardRecord var fr forwardRecord
@@ -152,7 +152,7 @@ func (h *Handler) getTunnelRecord(tunnelID int64) (*tunnelRecord, error) {
func (h *Handler) listForwardsByTunnel(tunnelID int64) ([]forwardRecord, error) { func (h *Handler) listForwardsByTunnel(tunnelID int64) ([]forwardRecord, error) {
rows, err := h.repo.DB().Query(` rows, err := h.repo.DB().Query(`
SELECT id, user_id, user_name, name, tunnel_id, remote_addr, strategy, status SELECT id, user_id, user_name, name, tunnel_id, remote_addr, COALESCE(strategy, 'fifo'), status
FROM forward FROM forward
WHERE tunnel_id = ? WHERE tunnel_id = ?
ORDER BY id ASC ORDER BY id ASC
+12 -2
View File
@@ -1356,7 +1356,7 @@ func parseIPLiteral(raw string) net.IP {
} }
if ip := net.ParseIP(value); ip != nil { if ip := net.ParseIP(value); ip != nil {
return ip return normalizeIPAddress(ip)
} }
host, _, err := net.SplitHostPort(value) host, _, err := net.SplitHostPort(value)
@@ -1368,7 +1368,17 @@ func parseIPLiteral(raw string) net.IP {
if host == "" { if host == "" {
return nil return nil
} }
return net.ParseIP(host) return normalizeIPAddress(net.ParseIP(host))
}
func normalizeIPAddress(ip net.IP) net.IP {
if ip == nil {
return nil
}
if v4 := ip.To4(); v4 != nil {
return v4
}
return ip.To16()
} }
func isTrustedProxyIP(ip net.IP) bool { func isTrustedProxyIP(ip net.IP) bool {
@@ -567,6 +567,13 @@ func TestAuthPeerAllowedIPs(t *testing.T) {
xff: "198.51.100.20, 172.20.0.3", xff: "198.51.100.20, 172.20.0.3",
wantAllowed: true, wantAllowed: true,
}, },
{
name: "ipv4-mapped proxy xff allowed",
allowedIPs: "198.51.100.20",
remoteAddr: "[::ffff:172.20.0.3]:34567",
xff: "198.51.100.20, 172.20.0.3",
wantAllowed: true,
},
{ {
name: "non whitelisted ip denied", name: "non whitelisted ip denied",
allowedIPs: "203.0.113.10", allowedIPs: "203.0.113.10",
@@ -251,7 +251,7 @@ func (h *Handler) pauseForwardRecords(forwards []forwardRecord, now int64) {
func (h *Handler) listActiveForwardsByUser(userID int64) ([]forwardRecord, error) { func (h *Handler) listActiveForwardsByUser(userID int64) ([]forwardRecord, error) {
rows, err := h.repo.DB().Query(` rows, err := h.repo.DB().Query(`
SELECT id, user_id, user_name, name, tunnel_id, remote_addr, strategy, status SELECT id, user_id, user_name, name, tunnel_id, remote_addr, COALESCE(strategy, 'fifo'), status
FROM forward FROM forward
WHERE user_id = ? AND status = 1 WHERE user_id = ? AND status = 1
ORDER BY id ASC ORDER BY id ASC
@@ -266,7 +266,7 @@ func (h *Handler) listActiveForwardsByUser(userID int64) ([]forwardRecord, error
func (h *Handler) listActiveForwardsByUserTunnel(userID int64, tunnelID int64) ([]forwardRecord, error) { func (h *Handler) listActiveForwardsByUserTunnel(userID int64, tunnelID int64) ([]forwardRecord, error) {
rows, err := h.repo.DB().Query(` rows, err := h.repo.DB().Query(`
SELECT id, user_id, user_name, name, tunnel_id, remote_addr, strategy, status SELECT id, user_id, user_name, name, tunnel_id, remote_addr, COALESCE(strategy, 'fifo'), status
FROM forward FROM forward
WHERE user_id = ? AND tunnel_id = ? AND status = 1 WHERE user_id = ? AND tunnel_id = ? AND status = 1
ORDER BY id ASC ORDER BY id ASC
@@ -2863,8 +2863,8 @@ func replaceTunnelChainsTx(tx *store.Tx, tunnelID int64, req map[string]interfac
if nodeID <= 0 { if nodeID <= 0 {
continue continue
} }
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, '1', ?, NULL, NULL, 0, ?)`, _, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, '1', ?, NULL, ?, 0, ?)`,
tunnelID, nodeID, defaultString(asString(n["protocol"]), "tls")) tunnelID, nodeID, defaultString(asString(n["strategy"]), "round"), defaultString(asString(n["protocol"]), "tls"))
if err != nil { if err != nil {
return err return err
} }
+34 -2
View File
@@ -719,7 +719,7 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
} }
rows, err := r.db.Query(` rows, err := r.db.Query(`
SELECT f.id, f.user_id, f.user_name, f.name, f.tunnel_id, COALESCE(t.name, ''), f.remote_addr, f.strategy, SELECT f.id, f.user_id, f.user_name, f.name, f.tunnel_id, COALESCE(t.name, ''), f.remote_addr, COALESCE(f.strategy, 'fifo'),
f.in_flow, f.out_flow, f.created_time, f.status, f.inx f.in_flow, f.out_flow, f.created_time, f.status, f.inx
FROM forward f FROM forward f
LEFT JOIN tunnel t ON t.id = f.tunnel_id LEFT JOIN tunnel t ON t.id = f.tunnel_id
@@ -1297,7 +1297,7 @@ func bootstrapSchema(db *store.DB, schemaSQL, seedSQL string) error {
return nil return nil
} }
const currentSchemaVersion = 1 const currentSchemaVersion = 2
func getSchemaVersion(db *store.DB) int { func getSchemaVersion(db *store.DB) int {
_, _ = db.Exec(`CREATE TABLE IF NOT EXISTS schema_version (version INTEGER NOT NULL DEFAULT 0)`) _, _ = db.Exec(`CREATE TABLE IF NOT EXISTS schema_version (version INTEGER NOT NULL DEFAULT 0)`)
@@ -1367,6 +1367,27 @@ func migrateSchema(db *store.DB) error {
} }
} }
normalizeStrategy := func(table, defaultValue string) error {
_, err := db.Exec(fmt.Sprintf("UPDATE %s SET strategy = ? WHERE strategy IS NULL", table), defaultValue)
if err != nil {
if isMissingTableError(db.Dialect(), err) {
return nil
}
return fmt.Errorf("normalize %s.strategy: %w", table, err)
}
return nil
}
if err := normalizeStrategy("forward", "fifo"); err != nil {
return err
}
if err := normalizeStrategy("chain_tunnel", "round"); err != nil {
return err
}
if err := normalizeStrategy("peer_share_runtime", "round"); err != nil {
return err
}
if db.Dialect() == store.DialectPostgres { if db.Dialect() == store.DialectPostgres {
if err := ensurePostgresIDDefaults(db); err != nil { if err := ensurePostgresIDDefaults(db); err != nil {
return err return err
@@ -1537,6 +1558,17 @@ func isMissingColumnError(dialect store.Dialect, err error) bool {
return strings.Contains(msg, "no such column") return strings.Contains(msg, "no such column")
} }
func isMissingTableError(dialect store.Dialect, err error) bool {
if err == nil {
return false
}
msg := strings.ToLower(err.Error())
if dialect == store.DialectPostgres {
return strings.Contains(msg, "relation") && strings.Contains(msg, "does not exist")
}
return strings.Contains(msg, "no such table")
}
func (r *Repository) CreatePeerShare(share *PeerShare) error { func (r *Repository) CreatePeerShare(share *PeerShare) error {
if r == nil || r.db == nil { if r == nil || r.db == nil {
return errors.New("repository not initialized") return errors.New("repository not initialized")
@@ -2,6 +2,7 @@ package contract_test
import ( import (
"bytes" "bytes"
"database/sql"
"encoding/json" "encoding/json"
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
@@ -144,6 +145,17 @@ func TestTunnelUpdateAssignsChainPortsContract(t *testing.T) {
if outPort <= 0 { if outPort <= 0 {
t.Fatalf("expected out node port to be assigned, got %d", outPort) t.Fatalf("expected out node port to be assigned, got %d", outPort)
} }
var entryStrategy sql.NullString
if err := repo.DB().QueryRow(`SELECT strategy FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 1 LIMIT 1`, tunnelID).Scan(&entryStrategy); err != nil {
t.Fatalf("query entry strategy: %v", err)
}
if !entryStrategy.Valid || strings.TrimSpace(entryStrategy.String) == "" {
t.Fatalf("expected entry strategy to be non-null and non-empty")
}
if entryStrategy.String != "round" {
t.Fatalf("expected entry strategy round, got %q", entryStrategy.String)
}
} }
func jsonInt(v int64) string { func jsonInt(v int64) string {