diff --git a/go-backend/internal/http/client/federation.go b/go-backend/internal/http/client/federation.go index 4fce731..8d5c0e9 100644 --- a/go-backend/internal/http/client/federation.go +++ b/go-backend/internal/http/client/federation.go @@ -69,6 +69,13 @@ type RuntimeReleaseRoleRequest struct { ResourceKey string `json:"resourceKey"` } +type RuntimeDiagnoseRequest struct { + IP string `json:"ip"` + Port int `json:"port"` + Count int `json:"count"` + Timeout int `json:"timeout"` +} + func NewFederationClient() *FederationClient { return &FederationClient{ client: &http.Client{ @@ -274,3 +281,46 @@ func (c *FederationClient) ReleaseRole(url, token, localDomain string, reqData R return nil } + +func (c *FederationClient) Diagnose(url, token, localDomain string, reqData RuntimeDiagnoseRequest) (map[string]interface{}, error) { + url = strings.TrimSuffix(url, "/") + bodyBytes, _ := json.Marshal(reqData) + req, err := http.NewRequest("POST", url+"/api/v1/federation/runtime/diagnose", strings.NewReader(string(bodyBytes))) + if err != nil { + return nil, err + } + req.Header.Set("Authorization", "Bearer "+token) + if localDomain != "" { + req.Header.Set("X-Panel-Domain", localDomain) + } + req.Header.Set("Content-Type", "application/json") + + resp, err := c.client.Do(req) + if err != nil { + return nil, err + } + defer resp.Body.Close() + + if resp.StatusCode != 200 { + body, _ := io.ReadAll(resp.Body) + return nil, fmt.Errorf("remote error %d: %s", resp.StatusCode, string(body)) + } + + var res struct { + Code int `json:"code"` + Msg string `json:"msg"` + Data map[string]interface{} `json:"data"` + } + if err := json.NewDecoder(resp.Body).Decode(&res); err != nil { + return nil, err + } + if res.Code != 0 { + return nil, fmt.Errorf("remote api error: %s", res.Msg) + } + + if res.Data == nil { + return nil, fmt.Errorf("remote api error: empty diagnosis payload") + } + + return res.Data, nil +} diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index c17cab3..0c3cf5a 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -11,6 +11,7 @@ import ( "strings" "time" + "go-backend/internal/http/client" "go-backend/internal/ws" ) @@ -791,7 +792,15 @@ func (h *Handler) appendPathDiagnosis(results *[]map[string]interface{}, nodeCac } item["nodeName"] = fromNode.Name - pingData, pingErr := h.tcpPingViaNode(fromNodeID, targetIP, targetPort) + var ( + pingData map[string]interface{} + pingErr error + ) + if fromNode.IsRemote == 1 { + pingData, pingErr = h.tcpPingViaRemoteNode(fromNode, targetIP, targetPort) + } else { + pingData, pingErr = h.tcpPingViaNode(fromNodeID, targetIP, targetPort) + } if pingErr != nil { item["success"] = false item["message"] = pingErr.Error() @@ -931,6 +940,25 @@ func (h *Handler) tcpPingViaNode(nodeID int64, ip string, port int) (map[string] return res.Data, nil } +func (h *Handler) tcpPingViaRemoteNode(node *nodeRecord, ip string, port int) (map[string]interface{}, error) { + if node == nil { + return nil, errors.New("节点不存在") + } + remoteURL := strings.TrimSpace(node.RemoteURL) + remoteToken := strings.TrimSpace(node.RemoteToken) + if remoteURL == "" || remoteToken == "" { + return nil, errors.New("远程节点缺少共享配置") + } + + fc := client.NewFederationClient() + return fc.Diagnose(remoteURL, remoteToken, h.federationLocalDomain(), client.RuntimeDiagnoseRequest{ + IP: strings.TrimSpace(ip), + Port: port, + Count: 4, + Timeout: 5000, + }) +} + func splitRemoteTargets(remoteAddr string) []string { parts := strings.Split(remoteAddr, ",") out := make([]string, 0, len(parts)) diff --git a/go-backend/internal/http/handler/federation.go b/go-backend/internal/http/handler/federation.go index 2a741a9..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 { @@ -65,6 +67,13 @@ type federationRuntimeReleaseRoleRequest struct { ResourceKey string `json:"resourceKey"` } +type federationRuntimeDiagnoseRequest struct { + IP string `json:"ip"` + Port int `json:"port"` + Count int `json:"count"` + Timeout int `json:"timeout"` +} + func (h *Handler) federationShareList(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPost { response.WriteJSON(w, response.ErrDefault("Invalid method")) @@ -116,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())) @@ -125,6 +140,10 @@ func (h *Handler) federationShareCreate(w http.ResponseWriter, r *http.Request) response.WriteJSON(w, response.ErrDefault("Node not found")) return } + if node.IsRemote == 1 { + response.WriteJSON(w, response.ErrDefault("Only local nodes can be shared")) + return + } now := time.Now().UnixMilli() token := randomToken(32) @@ -141,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 { @@ -283,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 == "" { @@ -731,6 +763,55 @@ func (h *Handler) federationRuntimeReleaseRole(w http.ResponseWriter, r *http.Re response.WriteJSON(w, response.OKEmpty()) } +func (h *Handler) federationRuntimeDiagnose(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("Invalid method")) + return + } + + token := extractBearerToken(r) + share, err := h.repo.GetPeerShareByToken(token) + if err != nil || share == nil { + response.WriteJSON(w, response.Err(401, "Unauthorized")) + return + } + + var req federationRuntimeDiagnoseRequest + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("Invalid JSON")) + return + } + + req.IP = strings.TrimSpace(req.IP) + if req.IP == "" || req.Port <= 0 || req.Port > 65535 { + response.WriteJSON(w, response.ErrDefault("Invalid target")) + return + } + if req.Count <= 0 { + req.Count = 4 + } + if req.Timeout <= 0 { + req.Timeout = 5000 + } + + res, err := h.sendNodeCommand(share.NodeID, "TcpPing", map[string]interface{}{ + "ip": req.IP, + "port": req.Port, + "count": req.Count, + "timeout": req.Timeout, + }, false, false) + if err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + if res.Data == nil { + response.WriteJSON(w, response.ErrDefault("Node did not return diagnosis data")) + return + } + + response.WriteJSON(w, response.OK(res.Data)) +} + func (h *Handler) pickPeerSharePort(share *sqlite.PeerShare, requestedPort int) (int, error) { if share == nil { return 0, fmt.Errorf("share not found") @@ -803,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 new file mode 100644 index 0000000..4e7b4da --- /dev/null +++ b/go-backend/internal/http/handler/federation_share_test.go @@ -0,0 +1,250 @@ +package handler + +import ( + "bytes" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "path/filepath" + "strings" + "testing" + "time" + + "go-backend/internal/http/response" + "go-backend/internal/store/sqlite" +) + +func TestFederationShareCreateRejectsRemoteNode(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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, "remote-share-node", "remote-share-secret", "10.10.10.1", "10.10.10.1", "", "20000-20010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "http://peer.example", "peer-token", `{"shareId":1}`) + if err != nil { + t.Fatalf("insert remote node: %v", err) + } + remoteNodeID, err := insertRes.LastInsertId() + if err != nil { + t.Fatalf("get remote node id: %v", err) + } + + body, err := json.Marshal(createPeerShareRequest{ + Name: "remote-node-share", + NodeID: remoteNodeID, + MaxBandwidth: 0, + ExpiryTime: 0, + PortRangeStart: 20000, + PortRangeEnd: 20010, + }) + 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 payload.Msg != "Only local nodes can be shared" { + t.Fatalf("expected rejection message %q, got %q", "Only local nodes can be shared", payload.Msg) + } + + var shareCount int + if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM peer_share WHERE node_id = ?`, remoteNodeID).Scan(&shareCount); err != nil { + t.Fatalf("query peer_share count: %v", err) + } + if shareCount != 0 { + 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/http/handler/handler.go b/go-backend/internal/http/handler/handler.go index d63ecb8..94451ee 100644 --- a/go-backend/internal/http/handler/handler.go +++ b/go-backend/internal/http/handler/handler.go @@ -161,6 +161,7 @@ func (h *Handler) Register(mux *http.ServeMux) { mux.HandleFunc("/api/v1/federation/runtime/reserve-port", h.authPeer(h.federationRuntimeReservePort)) mux.HandleFunc("/api/v1/federation/runtime/apply-role", h.authPeer(h.federationRuntimeApplyRole)) mux.HandleFunc("/api/v1/federation/runtime/release-role", h.authPeer(h.federationRuntimeReleaseRole)) + mux.HandleFunc("/api/v1/federation/runtime/diagnose", h.authPeer(h.federationRuntimeDiagnose)) mux.HandleFunc("/api/v1/federation/node/import", h.nodeImport) mux.HandleFunc("/flow/test", h.flowTest) diff --git a/go-backend/internal/http/middleware/auth.go b/go-backend/internal/http/middleware/auth.go index 93e8524..bf3aebf 100644 --- a/go-backend/internal/http/middleware/auth.go +++ b/go-backend/internal/http/middleware/auth.go @@ -91,6 +91,8 @@ func shouldSkip(path string) bool { return true case path == "/api/v1/federation/runtime/release-role": return true + case path == "/api/v1/federation/runtime/diagnose": + return true default: return false } diff --git a/go-backend/internal/store/sqlite/repository.go b/go-backend/internal/store/sqlite/repository.go index 4527584..9051f64 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": { "server_ip_v4": "VARCHAR(100)", @@ -1312,9 +1314,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 } @@ -1323,9 +1325,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 } @@ -1341,9 +1343,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 } @@ -1356,9 +1358,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 } @@ -1371,7 +1373,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 } @@ -1380,7 +1382,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/go-backend/tests/contract/diagnosis_contract_test.go b/go-backend/tests/contract/diagnosis_contract_test.go index cc3c76a..621b55d 100644 --- a/go-backend/tests/contract/diagnosis_contract_test.go +++ b/go-backend/tests/contract/diagnosis_contract_test.go @@ -8,6 +8,7 @@ import ( "path/filepath" "strconv" "strings" + "sync/atomic" "testing" "time" @@ -205,6 +206,173 @@ func TestDiagnosisChainCoverageContracts(t *testing.T) { }) } +func TestDiagnosisUsesFederationRuntimeForRemoteNodes(t *testing.T) { + secret := "contract-jwt-secret" + router, repo := setupDiagnosisContractRouter(t, secret) + now := time.Now().UnixMilli() + + remoteToken := "remote-diagnose-token" + var remoteDiagnoseCalls int32 + remoteServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/api/v1/federation/runtime/diagnose" { + http.NotFound(w, r) + return + } + if got := strings.TrimSpace(r.Header.Get("Authorization")); got != "Bearer "+remoteToken { + w.WriteHeader(http.StatusUnauthorized) + _ = json.NewEncoder(w).Encode(map[string]interface{}{"code": -1, "msg": "unauthorized"}) + return + } + + var req map[string]interface{} + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + w.WriteHeader(http.StatusBadRequest) + _ = json.NewEncoder(w).Encode(map[string]interface{}{"code": -1, "msg": "bad request"}) + return + } + if strings.TrimSpace(valueAsString(req["ip"])) != "10.50.0.30" { + w.WriteHeader(http.StatusBadRequest) + _ = json.NewEncoder(w).Encode(map[string]interface{}{"code": -1, "msg": "unexpected target ip"}) + return + } + if valueAsInt(req["port"]) != 30003 { + w.WriteHeader(http.StatusBadRequest) + _ = json.NewEncoder(w).Encode(map[string]interface{}{"code": -1, "msg": "unexpected target port"}) + return + } + + atomic.AddInt32(&remoteDiagnoseCalls, 1) + _ = json.NewEncoder(w).Encode(map[string]interface{}{ + "code": 0, + "msg": "success", + "data": map[string]interface{}{ + "success": true, + "averageTime": 12.5, + "packetLoss": 0, + "message": "remote tcp ok", + }, + }) + })) + defer remoteServer.Close() + + insertLocalNode := func(name, ip string) int64 { + res, 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) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, name, name+"-secret", ip, ip, "", "30000-30010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0) + if err != nil { + t.Fatalf("insert local node %s: %v", name, err) + } + id, err := res.LastInsertId() + if err != nil { + t.Fatalf("get local node id %s: %v", name, err) + } + return id + } + + insertRemoteNode := func(name, ip string) int64 { + res, 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(?, ?, ?, ?, ?, ?, ?, ?, 0, 0, 0, ?, ?, 1, ?, ?, ?, 1, ?, ?, ?) + `, name, name+"-secret", ip, "", "", "31000-31010", "", "", now, now, "[::]", "[::]", 1, remoteServer.URL, remoteToken, `{"shareId": 123}`) + if err != nil { + t.Fatalf("insert remote node %s: %v", name, err) + } + id, err := res.LastInsertId() + if err != nil { + t.Fatalf("get remote node id %s: %v", name, err) + } + return id + } + + entryNodeID := insertLocalNode("entry-local", "10.50.0.10") + remoteChainNodeID := insertRemoteNode("middle-remote", "10.50.0.20") + exitNodeID := insertLocalNode("exit-local", "10.50.0.30") + + tunnelRes, err := repo.DB().Exec(` + INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, "diagnose-remote-tunnel", 1.0, 2, "tls", 99999, now, now, 1, nil, 0) + if err != nil { + t.Fatalf("insert tunnel: %v", err) + } + tunnelID, err := tunnelRes.LastInsertId() + if err != nil { + t.Fatalf("get tunnel id: %v", err) + } + + if _, err := repo.DB().Exec(` + INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) + VALUES(?, 1, ?, 30001, 'round', 1, 'tls') + `, tunnelID, entryNodeID); err != nil { + t.Fatalf("insert entry chain: %v", err) + } + if _, err := repo.DB().Exec(` + INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) + VALUES(?, 2, ?, 30002, 'round', 1, 'tls') + `, tunnelID, remoteChainNodeID); err != nil { + t.Fatalf("insert middle chain: %v", err) + } + if _, err := repo.DB().Exec(` + INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) + VALUES(?, 3, ?, 30003, 'round', 1, 'tls') + `, tunnelID, exitNodeID); err != nil { + t.Fatalf("insert exit chain: %v", err) + } + + adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret) + if err != nil { + t.Fatalf("generate admin token: %v", err) + } + + req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/diagnose", bytes.NewBufferString(`{"tunnelId":`+strconv.FormatInt(tunnelID, 10)+`}`)) + req.Header.Set("Authorization", adminToken) + res := httptest.NewRecorder() + + router.ServeHTTP(res, req) + + var out response.R + if err := json.NewDecoder(res.Body).Decode(&out); err != nil { + t.Fatalf("decode response: %v", err) + } + if out.Code != 0 { + t.Fatalf("expected code 0, got %d (%s)", out.Code, out.Msg) + } + + payload, ok := out.Data.(map[string]interface{}) + if !ok { + t.Fatalf("expected object payload, got %T", out.Data) + } + results, ok := payload["results"].([]interface{}) + if !ok || len(results) == 0 { + t.Fatalf("expected non-empty results, got %v", payload["results"]) + } + + remoteStepFound := false + for _, raw := range results { + item, ok := raw.(map[string]interface{}) + if !ok { + continue + } + if valueAsInt(item["fromChainType"]) == 2 && valueAsInt(item["toChainType"]) == 3 { + remoteStepFound = true + if !valueAsBool(item["success"]) { + t.Fatalf("expected remote chain->exit diagnosis success, got item=%v", item) + } + if strings.TrimSpace(valueAsString(item["message"])) != "remote tcp ok" { + t.Fatalf("expected remote diagnosis message, got %q", valueAsString(item["message"])) + } + } + } + + if !remoteStepFound { + t.Fatalf("expected chain->exit diagnosis item for remote node") + } + if atomic.LoadInt32(&remoteDiagnoseCalls) == 0 { + t.Fatalf("expected federation runtime diagnose endpoint to be called") + } +} + func valueAsInt(v interface{}) int { switch n := v.(type) { case float64: @@ -223,6 +391,24 @@ func valueAsString(v interface{}) string { return s } +func valueAsBool(v interface{}) bool { + switch b := v.(type) { + case bool: + return b + case float64: + return b != 0 + case int: + return b != 0 + case int64: + return b != 0 + case string: + s := strings.TrimSpace(strings.ToLower(b)) + return s == "1" || s == "t" || s == "true" || s == "yes" || s == "y" + default: + return false + } +} + func setupDiagnosisContractRouter(t *testing.T, jwtSecret string) (http.Handler, *sqlite.Repository) { t.Helper() dbPath := filepath.Join(t.TempDir(), "diagnosis-contract.db") diff --git a/go-backend/tests/contract/federation_dual_panel_contract_test.go b/go-backend/tests/contract/federation_dual_panel_contract_test.go index 83f4782..b3bece9 100644 --- a/go-backend/tests/contract/federation_dual_panel_contract_test.go +++ b/go-backend/tests/contract/federation_dual_panel_contract_test.go @@ -15,6 +15,7 @@ import ( "github.com/gorilla/websocket" "go-backend/internal/auth" + "go-backend/internal/http/response" "go-backend/internal/security" "go-backend/internal/store/sqlite" ) @@ -152,6 +153,150 @@ func TestFederationDualPanelMiddleExitAutoPortContract(t *testing.T) { assertCount(t, providerRepo, `SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ?`, entryShareID, 0) } +func TestFederationDualPanelRemoteDiagnosisContract(t *testing.T) { + providerSecret := "provider-contract-jwt" + providerRouter, providerRepo := setupContractRouter(t, providerSecret) + providerServer := httptest.NewServer(providerRouter) + defer providerServer.Close() + + consumerSecret := "consumer-contract-jwt" + consumerRouter, consumerRepo := setupContractRouter(t, consumerSecret) + + consumerAdminToken, err := auth.GenerateToken(1, "consumer-admin", 0, consumerSecret) + if err != nil { + t.Fatalf("generate consumer admin token: %v", err) + } + + now := time.Now().UnixMilli() + providerEntryNodeID := insertContractNode(t, providerRepo, "provider-entry-dx", "203.0.113.11", "53000-53010", "provider-entry-dx-secret", 1) + providerMiddleNodeID := insertContractNode(t, providerRepo, "provider-middle-dx", "203.0.113.12", "54000-54010", "provider-middle-dx-secret", 1) + providerExitNodeID := insertContractNode(t, providerRepo, "provider-exit-dx", "203.0.113.13", "55000-55010", "provider-exit-dx-secret", 1) + + entryShareID := insertPeerShare(t, providerRepo, &sqlite.PeerShare{ + Name: "entry-share-dx", + NodeID: providerEntryNodeID, + Token: "share-entry-dx-token", + PortRangeStart: 53000, + PortRangeEnd: 53010, + IsActive: 1, + CreatedTime: now, + UpdatedTime: now, + }) + middleShareID := insertPeerShare(t, providerRepo, &sqlite.PeerShare{ + Name: "middle-share-dx", + NodeID: providerMiddleNodeID, + Token: "share-middle-dx-token", + PortRangeStart: 54000, + PortRangeEnd: 54010, + IsActive: 1, + CreatedTime: now, + UpdatedTime: now, + }) + exitShareID := insertPeerShare(t, providerRepo, &sqlite.PeerShare{ + Name: "exit-share-dx", + NodeID: providerExitNodeID, + Token: "share-exit-dx-token", + PortRangeStart: 55000, + PortRangeEnd: 55010, + IsActive: 1, + CreatedTime: now, + UpdatedTime: now, + }) + + importRemoteNodeForContract(t, consumerRouter, consumerAdminToken, providerServer.URL, "share-entry-dx-token") + importRemoteNodeForContract(t, consumerRouter, consumerAdminToken, providerServer.URL, "share-middle-dx-token") + importRemoteNodeForContract(t, consumerRouter, consumerAdminToken, providerServer.URL, "share-exit-dx-token") + + entryRemoteNodeID := queryRemoteNodeIDByToken(t, consumerRepo, "share-entry-dx-token") + middleRemoteNodeID := queryRemoteNodeIDByToken(t, consumerRepo, "share-middle-dx-token") + exitRemoteNodeID := queryRemoteNodeIDByToken(t, consumerRepo, "share-exit-dx-token") + + stopMiddle := startMockNodeSession(t, providerServer.URL, "provider-middle-dx-secret") + defer stopMiddle() + stopExit := startMockNodeSession(t, providerServer.URL, "provider-exit-dx-secret") + defer stopExit() + + createPayload := map[string]interface{}{ + "name": "dual-panel-diagnose-remote", + "type": 2, + "flow": 99999, + "status": 1, + "inNodeId": []map[string]interface{}{ + {"nodeId": entryRemoteNodeID, "protocol": "tls", "strategy": "round"}, + }, + "chainNodes": [][]map[string]interface{}{ + {{"nodeId": middleRemoteNodeID, "protocol": "tls", "strategy": "round"}}, + }, + "outNodeId": []map[string]interface{}{ + {"nodeId": exitRemoteNodeID, "protocol": "tls", "strategy": "round"}, + }, + } + body, err := json.Marshal(createPayload) + if err != nil { + t.Fatalf("marshal create payload: %v", err) + } + createReq := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/create", bytes.NewReader(body)) + createReq.Header.Set("Authorization", consumerAdminToken) + createReq.Header.Set("Content-Type", "application/json") + createRes := httptest.NewRecorder() + consumerRouter.ServeHTTP(createRes, createReq) + assertCode(t, createRes, 0) + + var tunnelID int64 + if err := consumerRepo.DB().QueryRow(`SELECT id FROM tunnel WHERE name = ? ORDER BY id DESC LIMIT 1`, "dual-panel-diagnose-remote").Scan(&tunnelID); err != nil { + t.Fatalf("query tunnel id: %v", err) + } + if tunnelID <= 0 { + t.Fatalf("invalid tunnel id") + } + + assertCount(t, providerRepo, `SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ? AND status = 1 AND applied = 1`, middleShareID, 1) + assertCount(t, providerRepo, `SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ? AND status = 1 AND applied = 1`, exitShareID, 1) + assertCount(t, providerRepo, `SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ?`, entryShareID, 0) + + diagnoseReq := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/diagnose", bytes.NewBufferString(fmt.Sprintf(`{"tunnelId":%d}`, tunnelID))) + diagnoseReq.Header.Set("Authorization", consumerAdminToken) + diagnoseRes := httptest.NewRecorder() + consumerRouter.ServeHTTP(diagnoseRes, diagnoseReq) + + var out response.R + if err := json.NewDecoder(diagnoseRes.Body).Decode(&out); err != nil { + t.Fatalf("decode diagnose response: %v", err) + } + if out.Code != 0 { + t.Fatalf("expected diagnose code 0, got %d (%s)", out.Code, out.Msg) + } + + payload, ok := out.Data.(map[string]interface{}) + if !ok { + t.Fatalf("expected map payload, got %T", out.Data) + } + results, ok := payload["results"].([]interface{}) + if !ok || len(results) == 0 { + t.Fatalf("expected non-empty results, got %v", payload["results"]) + } + + chainToExitFound := false + for _, raw := range results { + item, ok := raw.(map[string]interface{}) + if !ok { + continue + } + if valueAsInt(item["fromChainType"]) == 2 && valueAsInt(item["toChainType"]) == 3 { + chainToExitFound = true + if !valueAsBool(item["success"]) { + t.Fatalf("expected chain->exit diagnosis success, got item=%v", item) + } + if msg := strings.TrimSpace(valueAsString(item["message"])); msg != "mock tcp ok" { + t.Fatalf("expected remote diagnosis message 'mock tcp ok', got %q", msg) + } + } + } + if !chainToExitFound { + t.Fatalf("expected chain->exit diagnosis item in results") + } +} + func insertContractNode(t *testing.T, repo *sqlite.Repository, name, ip, portRange, secret string, status int) int64 { t.Helper() now := time.Now().UnixMilli() @@ -306,12 +451,21 @@ func startMockNodeSession(t *testing.T, baseURL string, nodeSecret string) func( } respType := fmt.Sprintf("%sResponse", cmd.Type) - respBytes, err := json.Marshal(map[string]interface{}{ + respPayload := map[string]interface{}{ "type": respType, "success": true, "message": "OK", "requestId": cmd.RequestID, - }) + } + if strings.EqualFold(strings.TrimSpace(cmd.Type), "TcpPing") { + respPayload["data"] = map[string]interface{}{ + "success": true, + "averageTime": 8.5, + "packetLoss": 0, + "message": "mock tcp ok", + } + } + respBytes, err := json.Marshal(respPayload) if err != nil { continue } @@ -324,3 +478,39 @@ func startMockNodeSession(t *testing.T, baseURL string, nodeSecret string) func( wg.Wait() } } + +func valueAsInt(v interface{}) int { + switch n := v.(type) { + case float64: + return int(n) + case int: + return n + case int64: + return int(n) + default: + return 0 + } +} + +func valueAsString(v interface{}) string { + s, _ := v.(string) + return s +} + +func valueAsBool(v interface{}) bool { + switch b := v.(type) { + case bool: + return b + case float64: + return b != 0 + case int: + return b != 0 + case int64: + return b != 0 + case string: + s := strings.TrimSpace(strings.ToLower(b)) + return s == "1" || s == "t" || s == "true" || s == "yes" || s == "y" + default: + return false + } +} 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 c78a9ec..e9afeae 100644 --- a/vite-frontend/src/pages/panel-sharing.tsx +++ b/vite-frontend/src/pages/panel-sharing.tsx @@ -1,4 +1,4 @@ -import { useState, useEffect } from "react"; +import { useState, useEffect, useCallback } from "react"; import { Button } from "@heroui/button"; import { Card, CardBody, CardHeader } from "@heroui/card"; import { Tabs, Tab } from "@heroui/tabs"; @@ -23,6 +23,7 @@ import { interface Node { id: number; name: string; + isRemote?: number; } interface PeerShare { @@ -35,6 +36,7 @@ interface PeerShare { portRangeEnd: number; isActive: number; allowedDomains?: string; + allowedIps?: string; } export default function PanelSharingPage() { @@ -56,6 +58,7 @@ export default function PanelSharingPage() { portRangeStart: 10000, portRangeEnd: 20000, allowedDomains: "", + allowedIps: "", }); const [importForm, setImportForm] = useState({ @@ -63,14 +66,7 @@ export default function PanelSharingPage() { token: "", }); - useEffect(() => { - if (selectedTab === "my-shares") { - loadShares(); - loadNodes(); - } - }, [selectedTab]); - - const loadShares = async () => { + const loadShares = useCallback(async () => { setLoading(true); try { const res = await getPeerShareList(); @@ -82,35 +78,60 @@ export default function PanelSharingPage() { } finally { setLoading(false); } - }; + }, []); - const loadNodes = async () => { + const loadNodes = useCallback(async () => { try { const res = await getNodeList(); if (res.code === 0) { - setNodes(res.data || []); + const localNodes: Node[] = (res.data || []).filter( + (node: Node) => (node?.isRemote ?? 0) !== 1, + ); + setNodes(localNodes); + setShareForm((prev) => { + if (!prev.nodeId) { + return prev; + } + const hasSelectedNode = localNodes.some( + (node: Node) => String(node.id) === prev.nodeId, + ); + return hasSelectedNode ? prev : { ...prev, nodeId: "" }; + }); } } catch { // ignore } - }; + }, []); + + useEffect(() => { + if (selectedTab === "my-shares") { + loadShares(); + loadNodes(); + } + }, [selectedTab, loadShares, loadNodes]); const handleCreateShare = async () => { if (!shareForm.name || !shareForm.nodeId) { toast.error("请填写必要信息"); return; } + const nodeId = parseInt(shareForm.nodeId, 10); + if (Number.isNaN(nodeId) || !nodes.some((node) => node.id === nodeId)) { + toast.error("仅可选择本地节点"); + return; + } try { const expiryTime = Date.now() + shareForm.expiryDays * 24 * 60 * 60 * 1000; const res = await createPeerShare({ name: shareForm.name, - nodeId: parseInt(shareForm.nodeId), + nodeId, maxBandwidth: shareForm.maxBandwidth * 1024 * 1024 * 1024, expiryTime: shareForm.expiryDays === 0 ? 0 : expiryTime, portRangeStart: shareForm.portRangeStart, portRangeEnd: shareForm.portRangeEnd, allowedDomains: shareForm.allowedDomains, + allowedIps: shareForm.allowedIps, }); if (res.code === 0) { toast.success("创建成功"); @@ -206,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()}

@@ -249,7 +271,7 @@ export default function PanelSharingPage() { /> setShareForm({ ...shareForm, allowedIps: e.target.value })} + /> @@ -321,4 +350,4 @@ export default function PanelSharingPage() {
); -} \ No newline at end of file +}