diff --git a/go-backend/internal/http/handler/federation.go b/go-backend/internal/http/handler/federation.go index 8c7ed0d..602dfc6 100644 --- a/go-backend/internal/http/handler/federation.go +++ b/go-backend/internal/http/handler/federation.go @@ -4,6 +4,7 @@ import ( "database/sql" "encoding/json" "fmt" + "net" "net/http" "strings" "time" @@ -27,6 +28,7 @@ type createPeerShareRequest struct { PortRangeStart int `json:"portRangeStart"` PortRangeEnd int `json:"portRangeEnd"` AllowedDomains string `json:"allowedDomains"` + AllowedIPs string `json:"allowedIps"` } type deletePeerShareRequest struct { @@ -123,6 +125,12 @@ func (h *Handler) federationShareCreate(w http.ResponseWriter, r *http.Request) return } + allowedIPs, err := normalizePeerShareAllowedIPs(req.AllowedIPs) + if err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + node, err := h.repo.GetNodeByID(req.NodeID) if err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) @@ -152,6 +160,7 @@ func (h *Handler) federationShareCreate(w http.ResponseWriter, r *http.Request) CreatedTime: now, UpdatedTime: now, AllowedDomains: req.AllowedDomains, + AllowedIPs: allowedIPs, } if err := h.repo.CreatePeerShare(share); err != nil { @@ -294,6 +303,18 @@ func (h *Handler) authPeer(next http.HandlerFunc) http.HandlerFunc { return } + if strings.TrimSpace(share.AllowedIPs) != "" { + clientIP := resolvePeerClientIP(r) + if clientIP == nil { + response.WriteJSON(w, response.Err(403, "Unable to determine client IP")) + return + } + if !isPeerIPAllowed(clientIP, share.AllowedIPs) { + response.WriteJSON(w, response.Err(403, "IP not allowed")) + return + } + } + if share.AllowedDomains != "" { clientDomain := r.Header.Get("X-Panel-Domain") if clientDomain == "" { @@ -863,3 +884,130 @@ func extractBearerToken(r *http.Request) string { } return "" } + +func normalizePeerShareAllowedIPs(raw string) (string, error) { + raw = strings.TrimSpace(raw) + if raw == "" { + return "", nil + } + + parts := strings.Split(raw, ",") + normalized := make([]string, 0, len(parts)) + seen := make(map[string]struct{}, len(parts)) + + for _, part := range parts { + item := strings.TrimSpace(part) + if item == "" { + continue + } + + if strings.Contains(item, "/") { + _, network, err := net.ParseCIDR(item) + if err != nil { + return "", fmt.Errorf("Invalid allowed IP or CIDR: %s", item) + } + item = network.String() + } else { + ip := parseIPLiteral(item) + if ip == nil { + return "", fmt.Errorf("Invalid allowed IP or CIDR: %s", item) + } + item = ip.String() + } + + if _, exists := seen[item]; exists { + continue + } + seen[item] = struct{}{} + normalized = append(normalized, item) + } + + return strings.Join(normalized, ","), nil +} + +func resolvePeerClientIP(r *http.Request) net.IP { + if r == nil { + return nil + } + + remoteIP := parseIPLiteral(r.RemoteAddr) + if isTrustedProxyIP(remoteIP) { + if ip := parseForwardedFor(r.Header.Get("X-Forwarded-For")); ip != nil { + return ip + } + if ip := parseIPLiteral(r.Header.Get("X-Real-IP")); ip != nil { + return ip + } + } + + return remoteIP +} + +func parseForwardedFor(raw string) net.IP { + for _, part := range strings.Split(raw, ",") { + if ip := parseIPLiteral(part); ip != nil { + return ip + } + } + return nil +} + +func parseIPLiteral(raw string) net.IP { + value := strings.Trim(strings.TrimSpace(raw), "\"") + if value == "" { + return nil + } + + if ip := net.ParseIP(value); ip != nil { + return ip + } + + host, _, err := net.SplitHostPort(value) + if err != nil { + return nil + } + + host = strings.Trim(strings.TrimSpace(host), "[]") + if host == "" { + return nil + } + return net.ParseIP(host) +} + +func isTrustedProxyIP(ip net.IP) bool { + if ip == nil { + return false + } + return ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() +} + +func isPeerIPAllowed(clientIP net.IP, whitelist string) bool { + if clientIP == nil { + return false + } + + for _, part := range strings.Split(whitelist, ",") { + entry := strings.TrimSpace(part) + if entry == "" { + continue + } + + if strings.Contains(entry, "/") { + _, network, err := net.ParseCIDR(entry) + if err != nil { + continue + } + if network.Contains(clientIP) { + return true + } + continue + } + + allowedIP := parseIPLiteral(entry) + if allowedIP != nil && allowedIP.Equal(clientIP) { + return true + } + } + + return false +} diff --git a/go-backend/internal/http/handler/federation_share_test.go b/go-backend/internal/http/handler/federation_share_test.go index f0b4746..4e7b4da 100644 --- a/go-backend/internal/http/handler/federation_share_test.go +++ b/go-backend/internal/http/handler/federation_share_test.go @@ -3,9 +3,11 @@ package handler import ( "bytes" "encoding/json" + "fmt" "net/http" "net/http/httptest" "path/filepath" + "strings" "testing" "time" @@ -76,3 +78,173 @@ func TestFederationShareCreateRejectsRemoteNode(t *testing.T) { t.Fatalf("expected no share rows for remote node, got %d", shareCount) } } + +func TestFederationShareCreateRejectsInvalidAllowedIPs(t *testing.T) { + repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db")) + if err != nil { + t.Fatalf("open sqlite: %v", err) + } + t.Cleanup(func() { _ = repo.Close() }) + + h := New(repo, "test-jwt-secret") + now := time.Now().UnixMilli() + + insertRes, err := repo.DB().Exec(` + INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, "local-share-node", "local-share-secret", "10.20.30.40", "10.20.30.40", "", "21000-21010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 0, "", "", "") + if err != nil { + t.Fatalf("insert local node: %v", err) + } + localNodeID, err := insertRes.LastInsertId() + if err != nil { + t.Fatalf("get local node id: %v", err) + } + + body, err := json.Marshal(createPeerShareRequest{ + Name: "local-node-share", + NodeID: localNodeID, + MaxBandwidth: 0, + ExpiryTime: 0, + PortRangeStart: 21000, + PortRangeEnd: 21010, + AllowedIPs: "bad-ip-entry", + }) + if err != nil { + t.Fatalf("marshal request: %v", err) + } + + req := httptest.NewRequest(http.MethodPost, "/api/v1/federation/share/create", bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + res := httptest.NewRecorder() + + h.federationShareCreate(res, req) + + if res.Code != http.StatusOK { + t.Fatalf("expected status %d, got %d", http.StatusOK, res.Code) + } + + var payload response.R + if err := json.NewDecoder(res.Body).Decode(&payload); err != nil { + t.Fatalf("decode response: %v", err) + } + if payload.Code != -1 { + t.Fatalf("expected response code -1, got %d", payload.Code) + } + if !strings.Contains(payload.Msg, "Invalid allowed IP or CIDR") { + t.Fatalf("expected invalid IP message, got %q", payload.Msg) + } + + var shareCount int + if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM peer_share WHERE node_id = ?`, localNodeID).Scan(&shareCount); err != nil { + t.Fatalf("query peer_share count: %v", err) + } + if shareCount != 0 { + t.Fatalf("expected no share rows for node, got %d", shareCount) + } +} + +func TestAuthPeerAllowedIPs(t *testing.T) { + repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db")) + if err != nil { + t.Fatalf("open sqlite: %v", err) + } + t.Cleanup(func() { _ = repo.Close() }) + + h := New(repo, "test-jwt-secret") + now := time.Now().UnixMilli() + + tests := []struct { + name string + allowedIPs string + remoteAddr string + xff string + wantAllowed bool + }{ + { + name: "exact ip allowed", + allowedIPs: "203.0.113.10", + remoteAddr: "203.0.113.10:23456", + wantAllowed: true, + }, + { + name: "cidr allowed", + allowedIPs: "203.0.113.0/24", + remoteAddr: "203.0.113.11:23456", + wantAllowed: true, + }, + { + name: "trusted proxy xff allowed", + allowedIPs: "198.51.100.20", + remoteAddr: "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", + remoteAddr: "203.0.113.99:23456", + wantAllowed: false, + }, + } + + for idx, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + token := fmt.Sprintf("share-token-%d", idx) + if err := repo.CreatePeerShare(&sqlite.PeerShare{ + Name: "share-" + tt.name, + NodeID: 1, + Token: token, + PortRangeStart: 10000, + PortRangeEnd: 10010, + IsActive: 1, + CreatedTime: now, + UpdatedTime: now, + AllowedIPs: tt.allowedIPs, + }); err != nil { + t.Fatalf("create peer share: %v", err) + } + + nextCalled := false + wrapped := h.authPeer(func(w http.ResponseWriter, r *http.Request) { + nextCalled = true + response.WriteJSON(w, response.OKEmpty()) + }) + + req := httptest.NewRequest(http.MethodPost, "/api/v1/federation/connect", nil) + req.Header.Set("Authorization", "Bearer "+token) + if tt.xff != "" { + req.Header.Set("X-Forwarded-For", tt.xff) + } + req.RemoteAddr = tt.remoteAddr + + res := httptest.NewRecorder() + wrapped(res, req) + + var payload response.R + if err := json.NewDecoder(res.Body).Decode(&payload); err != nil { + t.Fatalf("decode response: %v", err) + } + + if tt.wantAllowed { + if !nextCalled { + t.Fatalf("expected next handler to be called") + } + if payload.Code != 0 { + t.Fatalf("expected code 0, got %d (%s)", payload.Code, payload.Msg) + } + return + } + + if nextCalled { + t.Fatalf("expected next handler not to be called") + } + if payload.Code != 403 { + t.Fatalf("expected code 403, got %d (%s)", payload.Code, payload.Msg) + } + if payload.Msg != "IP not allowed" { + t.Fatalf("expected IP rejection message, got %q", payload.Msg) + } + }) + } +} diff --git a/go-backend/internal/store/sqlite/repository.go b/go-backend/internal/store/sqlite/repository.go index f4e5e49..5c43c89 100644 --- a/go-backend/internal/store/sqlite/repository.go +++ b/go-backend/internal/store/sqlite/repository.go @@ -122,6 +122,7 @@ type PeerShare struct { CreatedTime int64 `json:"createdTime"` UpdatedTime int64 `json:"updatedTime"` AllowedDomains string `json:"allowedDomains"` + AllowedIPs string `json:"allowedIps"` } type PeerShareRuntime struct { @@ -1278,6 +1279,7 @@ func migrateSchema(db *sql.DB) error { columnsByTable := map[string]map[string]string{ "peer_share": { "allowed_domains": "TEXT DEFAULT ''", + "allowed_ips": "TEXT DEFAULT ''", }, "node": { "inx": "INTEGER NOT NULL DEFAULT 0", @@ -1310,9 +1312,9 @@ func (r *Repository) CreatePeerShare(share *PeerShare) error { return errors.New("repository not initialized") } _, err := r.db.Exec(` - INSERT INTO peer_share(name, node_id, token, max_bandwidth, expiry_time, port_range_start, port_range_end, current_flow, is_active, created_time, updated_time, allowed_domains) - VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, share.Name, share.NodeID, share.Token, share.MaxBandwidth, share.ExpiryTime, share.PortRangeStart, share.PortRangeEnd, share.CurrentFlow, share.IsActive, share.CreatedTime, share.UpdatedTime, share.AllowedDomains) + INSERT INTO peer_share(name, node_id, token, max_bandwidth, expiry_time, port_range_start, port_range_end, current_flow, is_active, created_time, updated_time, allowed_domains, allowed_ips) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, share.Name, share.NodeID, share.Token, share.MaxBandwidth, share.ExpiryTime, share.PortRangeStart, share.PortRangeEnd, share.CurrentFlow, share.IsActive, share.CreatedTime, share.UpdatedTime, share.AllowedDomains, share.AllowedIPs) return err } @@ -1321,9 +1323,9 @@ func (r *Repository) UpdatePeerShare(share *PeerShare) error { return errors.New("repository not initialized") } _, err := r.db.Exec(` - UPDATE peer_share SET name=?, max_bandwidth=?, expiry_time=?, port_range_start=?, port_range_end=?, is_active=?, updated_time=?, allowed_domains=? + UPDATE peer_share SET name=?, max_bandwidth=?, expiry_time=?, port_range_start=?, port_range_end=?, is_active=?, updated_time=?, allowed_domains=?, allowed_ips=? WHERE id=? - `, share.Name, share.MaxBandwidth, share.ExpiryTime, share.PortRangeStart, share.PortRangeEnd, share.IsActive, share.UpdatedTime, share.AllowedDomains, share.ID) + `, share.Name, share.MaxBandwidth, share.ExpiryTime, share.PortRangeStart, share.PortRangeEnd, share.IsActive, share.UpdatedTime, share.AllowedDomains, share.AllowedIPs, share.ID) return err } @@ -1339,9 +1341,9 @@ func (r *Repository) GetPeerShare(id int64) (*PeerShare, error) { if r == nil || r.db == nil { return nil, errors.New("repository not initialized") } - row := r.db.QueryRow(`SELECT id, name, node_id, token, max_bandwidth, expiry_time, port_range_start, port_range_end, current_flow, is_active, created_time, updated_time, allowed_domains FROM peer_share WHERE id = ?`, id) + row := r.db.QueryRow(`SELECT id, name, node_id, token, max_bandwidth, expiry_time, port_range_start, port_range_end, current_flow, is_active, created_time, updated_time, allowed_domains, allowed_ips FROM peer_share WHERE id = ?`, id) var s PeerShare - if err := row.Scan(&s.ID, &s.Name, &s.NodeID, &s.Token, &s.MaxBandwidth, &s.ExpiryTime, &s.PortRangeStart, &s.PortRangeEnd, &s.CurrentFlow, &s.IsActive, &s.CreatedTime, &s.UpdatedTime, &s.AllowedDomains); err != nil { + if err := row.Scan(&s.ID, &s.Name, &s.NodeID, &s.Token, &s.MaxBandwidth, &s.ExpiryTime, &s.PortRangeStart, &s.PortRangeEnd, &s.CurrentFlow, &s.IsActive, &s.CreatedTime, &s.UpdatedTime, &s.AllowedDomains, &s.AllowedIPs); err != nil { if errors.Is(err, sql.ErrNoRows) { return nil, nil } @@ -1354,9 +1356,9 @@ func (r *Repository) GetPeerShareByToken(token string) (*PeerShare, error) { if r == nil || r.db == nil { return nil, errors.New("repository not initialized") } - row := r.db.QueryRow(`SELECT id, name, node_id, token, max_bandwidth, expiry_time, port_range_start, port_range_end, current_flow, is_active, created_time, updated_time, allowed_domains FROM peer_share WHERE token = ?`, token) + row := r.db.QueryRow(`SELECT id, name, node_id, token, max_bandwidth, expiry_time, port_range_start, port_range_end, current_flow, is_active, created_time, updated_time, allowed_domains, allowed_ips FROM peer_share WHERE token = ?`, token) var s PeerShare - if err := row.Scan(&s.ID, &s.Name, &s.NodeID, &s.Token, &s.MaxBandwidth, &s.ExpiryTime, &s.PortRangeStart, &s.PortRangeEnd, &s.CurrentFlow, &s.IsActive, &s.CreatedTime, &s.UpdatedTime, &s.AllowedDomains); err != nil { + if err := row.Scan(&s.ID, &s.Name, &s.NodeID, &s.Token, &s.MaxBandwidth, &s.ExpiryTime, &s.PortRangeStart, &s.PortRangeEnd, &s.CurrentFlow, &s.IsActive, &s.CreatedTime, &s.UpdatedTime, &s.AllowedDomains, &s.AllowedIPs); err != nil { if errors.Is(err, sql.ErrNoRows) { return nil, nil } @@ -1369,7 +1371,7 @@ func (r *Repository) ListPeerShares() ([]PeerShare, error) { if r == nil || r.db == nil { return nil, errors.New("repository not initialized") } - rows, err := r.db.Query(`SELECT id, name, node_id, token, max_bandwidth, expiry_time, port_range_start, port_range_end, current_flow, is_active, created_time, updated_time, allowed_domains FROM peer_share ORDER BY id DESC`) + rows, err := r.db.Query(`SELECT id, name, node_id, token, max_bandwidth, expiry_time, port_range_start, port_range_end, current_flow, is_active, created_time, updated_time, allowed_domains, allowed_ips FROM peer_share ORDER BY id DESC`) if err != nil { return nil, err } @@ -1378,7 +1380,7 @@ func (r *Repository) ListPeerShares() ([]PeerShare, error) { var shares []PeerShare for rows.Next() { var s PeerShare - if err := rows.Scan(&s.ID, &s.Name, &s.NodeID, &s.Token, &s.MaxBandwidth, &s.ExpiryTime, &s.PortRangeStart, &s.PortRangeEnd, &s.CurrentFlow, &s.IsActive, &s.CreatedTime, &s.UpdatedTime, &s.AllowedDomains); err != nil { + if err := rows.Scan(&s.ID, &s.Name, &s.NodeID, &s.Token, &s.MaxBandwidth, &s.ExpiryTime, &s.PortRangeStart, &s.PortRangeEnd, &s.CurrentFlow, &s.IsActive, &s.CreatedTime, &s.UpdatedTime, &s.AllowedDomains, &s.AllowedIPs); err != nil { return nil, err } shares = append(shares, s) diff --git a/go-backend/internal/store/sqlite/sql/schema.sql b/go-backend/internal/store/sqlite/sql/schema.sql index 74fcf0e..67a6a39 100644 --- a/go-backend/internal/store/sqlite/sql/schema.sql +++ b/go-backend/internal/store/sqlite/sql/schema.sql @@ -199,7 +199,8 @@ CREATE TABLE IF NOT EXISTS peer_share ( is_active INTEGER DEFAULT 1, created_time INTEGER NOT NULL, updated_time INTEGER NOT NULL, - allowed_domains TEXT DEFAULT '' + allowed_domains TEXT DEFAULT '', + allowed_ips TEXT DEFAULT '' ); CREATE TABLE IF NOT EXISTS peer_share_runtime ( diff --git a/vite-frontend/src/api/index.ts b/vite-frontend/src/api/index.ts index c568221..effec4e 100644 --- a/vite-frontend/src/api/index.ts +++ b/vite-frontend/src/api/index.ts @@ -198,6 +198,7 @@ export const createPeerShare = (data: { portRangeStart?: number; portRangeEnd?: number; allowedDomains?: string; + allowedIps?: string; }) => Network.post("/federation/share/create", data); export const deletePeerShare = (id: number) => Network.post("/federation/share/delete", { id }); diff --git a/vite-frontend/src/pages/panel-sharing.tsx b/vite-frontend/src/pages/panel-sharing.tsx index e179648..e9afeae 100644 --- a/vite-frontend/src/pages/panel-sharing.tsx +++ b/vite-frontend/src/pages/panel-sharing.tsx @@ -36,6 +36,7 @@ interface PeerShare { portRangeEnd: number; isActive: number; allowedDomains?: string; + allowedIps?: string; } export default function PanelSharingPage() { @@ -57,6 +58,7 @@ export default function PanelSharingPage() { portRangeStart: 10000, portRangeEnd: 20000, allowedDomains: "", + allowedIps: "", }); const [importForm, setImportForm] = useState({ @@ -129,6 +131,7 @@ export default function PanelSharingPage() { portRangeStart: shareForm.portRangeStart, portRangeEnd: shareForm.portRangeEnd, allowedDomains: shareForm.allowedDomains, + allowedIps: shareForm.allowedIps, }); if (res.code === 0) { toast.success("创建成功"); @@ -224,6 +227,7 @@ export default function PanelSharingPage() {

端口范围: {share.portRangeStart} - {share.portRangeEnd}

{share.allowedDomains &&

允许域名: {share.allowedDomains}

} + {share.allowedIps &&

允许API IP: {share.allowedIps}

}

过期时间: {share.expiryTime === 0 ? "永久" : new Date(share.expiryTime).toLocaleDateString()}

@@ -305,6 +309,13 @@ export default function PanelSharingPage() { value={shareForm.allowedDomains} onChange={(e) => setShareForm({ ...shareForm, allowedDomains: e.target.value })} /> + setShareForm({ ...shareForm, allowedIps: e.target.value })} + />