mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-08 10:46:37 +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) {
|
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
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
Reference in New Issue
Block a user