feat(federation): add import API IP whitelist controls

This commit is contained in:
sagit
2026-02-10 06:40:46 +00:00
parent 00079ac7af
commit b205b47414
6 changed files with 347 additions and 12 deletions
@@ -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
}
@@ -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)
}
})
}
}
+13 -11
View File
@@ -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)
@@ -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 (
+1
View File
@@ -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 });
+11
View File
@@ -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() {
<CardBody className="text-sm space-y-2">
<p>端口范围: {share.portRangeStart} - {share.portRangeEnd}</p>
{share.allowedDomains && <p>允许域名: {share.allowedDomains}</p>}
{share.allowedIps && <p>允许API IP: {share.allowedIps}</p>}
<p>过期时间: {share.expiryTime === 0 ? "永久" : new Date(share.expiryTime).toLocaleDateString()}</p>
<div className="flex gap-2">
<Input readOnly size="sm" value={share.token} />
@@ -305,6 +309,13 @@ export default function PanelSharingPage() {
value={shareForm.allowedDomains}
onChange={(e) => setShareForm({ ...shareForm, allowedDomains: e.target.value })}
/>
<Input
label="允许的API IP (可选)"
placeholder="203.0.113.10, 2001:db8::10, 198.51.100.0/24"
description="仅白名单IP可导入此分享,支持IPv4/IPv6/CIDR,多个用逗号分隔"
value={shareForm.allowedIps}
onChange={(e) => setShareForm({ ...shareForm, allowedIps: e.target.value })}
/>
</ModalBody>
<ModalFooter>
<Button onPress={() => setCreateShareOpen(false)}>取消</Button>