mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-29 16:06:36 +08:00
fix(backend): normalize strategy data and proxy ip parsing
This commit is contained in:
@@ -114,7 +114,7 @@ func (h *Handler) ensureTunnelPermission(userID int64, roleID int, tunnelID int6
|
||||
|
||||
func (h *Handler) getForwardRecord(forwardID int64) (*forwardRecord, error) {
|
||||
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
|
||||
`, forwardID)
|
||||
var fr forwardRecord
|
||||
@@ -152,7 +152,7 @@ func (h *Handler) getTunnelRecord(tunnelID int64) (*tunnelRecord, error) {
|
||||
|
||||
func (h *Handler) listForwardsByTunnel(tunnelID int64) ([]forwardRecord, error) {
|
||||
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
|
||||
WHERE tunnel_id = ?
|
||||
ORDER BY id ASC
|
||||
|
||||
@@ -1356,7 +1356,7 @@ func parseIPLiteral(raw string) net.IP {
|
||||
}
|
||||
|
||||
if ip := net.ParseIP(value); ip != nil {
|
||||
return ip
|
||||
return normalizeIPAddress(ip)
|
||||
}
|
||||
|
||||
host, _, err := net.SplitHostPort(value)
|
||||
@@ -1368,7 +1368,17 @@ func parseIPLiteral(raw string) net.IP {
|
||||
if host == "" {
|
||||
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 {
|
||||
|
||||
@@ -567,6 +567,13 @@ func TestAuthPeerAllowedIPs(t *testing.T) {
|
||||
xff: "198.51.100.20, 172.20.0.3",
|
||||
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",
|
||||
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) {
|
||||
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
|
||||
WHERE user_id = ? AND status = 1
|
||||
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) {
|
||||
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
|
||||
WHERE user_id = ? AND tunnel_id = ? AND status = 1
|
||||
ORDER BY id ASC
|
||||
|
||||
@@ -2863,8 +2863,8 @@ func replaceTunnelChainsTx(tx *store.Tx, tunnelID int64, req map[string]interfac
|
||||
if nodeID <= 0 {
|
||||
continue
|
||||
}
|
||||
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, '1', ?, NULL, NULL, 0, ?)`,
|
||||
tunnelID, nodeID, defaultString(asString(n["protocol"]), "tls"))
|
||||
_, 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["strategy"]), "round"), defaultString(asString(n["protocol"]), "tls"))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -719,7 +719,7 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
|
||||
}
|
||||
|
||||
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
|
||||
FROM forward f
|
||||
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
|
||||
}
|
||||
|
||||
const currentSchemaVersion = 1
|
||||
const currentSchemaVersion = 2
|
||||
|
||||
func getSchemaVersion(db *store.DB) int {
|
||||
_, _ = 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 err := ensurePostgresIDDefaults(db); err != nil {
|
||||
return err
|
||||
@@ -1537,6 +1558,17 @@ func isMissingColumnError(dialect store.Dialect, err error) bool {
|
||||
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 {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
|
||||
@@ -2,6 +2,7 @@ package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -144,6 +145,17 @@ func TestTunnelUpdateAssignsChainPortsContract(t *testing.T) {
|
||||
if outPort <= 0 {
|
||||
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 {
|
||||
|
||||
Reference in New Issue
Block a user