diff --git a/go-backend/internal/http/handler/federation_runtime_test.go b/go-backend/internal/http/handler/federation_runtime_test.go new file mode 100644 index 0000000..8aa7cc8 --- /dev/null +++ b/go-backend/internal/http/handler/federation_runtime_test.go @@ -0,0 +1,148 @@ +package handler + +import ( + "path/filepath" + "testing" + "time" + + "go-backend/internal/store/sqlite" +) + +func TestPickPeerSharePortUsesRuntimeReservations(t *testing.T) { + repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db")) + if err != nil { + t.Fatalf("open repo: %v", err) + } + defer repo.Close() + + h := &Handler{repo: repo} + now := time.Now().UnixMilli() + + if _, err := repo.DB().Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, ?, ?, ?, ?, ?, ?)`, 1, 2, 1, 3000, "round", 1, "tls"); err != nil { + t.Fatalf("insert chain_tunnel: %v", err) + } + if _, err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, 1, 1, 3001); err != nil { + t.Fatalf("insert forward_port: %v", err) + } + if _, err := repo.DB().Exec(` + INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, 77, 1, "res-1", "rk-1", "b-1", "exit", "", "fed_svc_1", "tls", "round", 3002, "", 1, 1, now, now); err != nil { + t.Fatalf("insert peer_share_runtime: %v", err) + } + + share := &sqlite.PeerShare{ + ID: 77, + NodeID: 1, + PortRangeStart: 3000, + PortRangeEnd: 3004, + } + + port, err := h.pickPeerSharePort(share, 0) + if err != nil { + t.Fatalf("pick auto port: %v", err) + } + if port != 3003 { + t.Fatalf("expected port 3003, got %d", port) + } + + if _, err := h.pickPeerSharePort(share, 3001); err == nil { + t.Fatalf("expected requested busy port to fail") + } +} + +func TestApplyTunnelRuntimeSkipsRemoteNodes(t *testing.T) { + h := &Handler{} + state := &tunnelCreateState{ + TunnelID: 1, + Type: 2, + InNodes: []tunnelRuntimeNode{ + {NodeID: 11, ChainType: 1, Protocol: "tls"}, + }, + ChainHops: [][]tunnelRuntimeNode{ + { + {NodeID: 12, ChainType: 2, Inx: 1, Port: 41000, Protocol: "tls", Strategy: "round"}, + }, + }, + OutNodes: []tunnelRuntimeNode{ + {NodeID: 13, ChainType: 3, Port: 42000, Protocol: "tls", Strategy: "round"}, + }, + Nodes: map[int64]*nodeRecord{ + 11: {ID: 11, Name: "remote-in", IsRemote: 1}, + 12: {ID: 12, Name: "remote-chain", IsRemote: 1}, + 13: {ID: 13, Name: "remote-out", IsRemote: 1}, + }, + } + + chains, services, err := h.applyTunnelRuntime(state) + if err != nil { + t.Fatalf("apply runtime: %v", err) + } + if len(chains) != 0 { + t.Fatalf("expected no local chains created, got %d", len(chains)) + } + if len(services) != 0 { + t.Fatalf("expected no local services created, got %d", len(services)) + } +} + +func TestPrepareTunnelCreateStateRemoteAutoPortDefersToFederation(t *testing.T) { + repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db")) + if err != nil { + t.Fatalf("open repo: %v", err) + } + defer repo.Close() + + h := &Handler{repo: repo} + now := time.Now().UnixMilli() + + insertNode := func(name string, status int, portRange string, isRemote int) int64 { + res, execErr := 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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, name, name+"-secret", "10.0.0.1", "10.0.0.1", "", portRange, "", "v1", 1, 1, 1, now, now, status, "[::]", "[::]", 0, isRemote, "http://peer", "peer-token", `{"shareId":1}`) + if execErr != nil { + t.Fatalf("insert node %s: %v", name, execErr) + } + id, idErr := res.LastInsertId() + if idErr != nil { + t.Fatalf("node id %s: %v", name, idErr) + } + return id + } + + entryID := insertNode("entry", 1, "31000-31010", 0) + remoteOutID := insertNode("remote-out", 1, "30000", 1) + + if _, err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, 1, remoteOutID, 30000); err != nil { + t.Fatalf("insert forward_port: %v", err) + } + + tx, err := repo.DB().Begin() + if err != nil { + t.Fatalf("begin tx: %v", err) + } + defer tx.Rollback() + + req := map[string]interface{}{ + "name": "test-tunnel", + "inNodeId": []interface{}{ + map[string]interface{}{"nodeId": float64(entryID), "protocol": "tls", "strategy": "round"}, + }, + "outNodeId": []interface{}{ + map[string]interface{}{"nodeId": float64(remoteOutID), "protocol": "tls", "strategy": "round", "port": float64(0)}, + }, + "chainNodes": []interface{}{}, + } + + state, err := h.prepareTunnelCreateState(tx, req, 2, 0) + if err != nil { + t.Fatalf("prepare state should not fail for remote auto-port: %v", err) + } + if len(state.OutNodes) != 1 { + t.Fatalf("expected 1 out node, got %d", len(state.OutNodes)) + } + if state.OutNodes[0].Port != 0 { + t.Fatalf("expected remote out port to remain 0 before federation reserve, got %d", state.OutNodes[0].Port) + } +} diff --git a/go-backend/tests/contract/federation_dual_panel_contract_test.go b/go-backend/tests/contract/federation_dual_panel_contract_test.go new file mode 100644 index 0000000..83f4782 --- /dev/null +++ b/go-backend/tests/contract/federation_dual_panel_contract_test.go @@ -0,0 +1,326 @@ +package contract_test + +import ( + "bytes" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "sync" + "testing" + "time" + + "github.com/gorilla/websocket" + + "go-backend/internal/auth" + "go-backend/internal/security" + "go-backend/internal/store/sqlite" +) + +func TestFederationDualPanelMiddleExitAutoPortContract(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", "198.51.100.11", "43000-43010", "provider-entry-secret", 1) + providerMiddleNodeID := insertContractNode(t, providerRepo, "provider-middle", "198.51.100.12", "44000-44010", "provider-middle-secret", 1) + providerExitNodeID := insertContractNode(t, providerRepo, "provider-exit", "198.51.100.13", "45000-45010", "provider-exit-secret", 1) + + entryShareID := insertPeerShare(t, providerRepo, &sqlite.PeerShare{ + Name: "entry-share", + NodeID: providerEntryNodeID, + Token: "share-entry-token", + PortRangeStart: 43000, + PortRangeEnd: 43010, + IsActive: 1, + CreatedTime: now, + UpdatedTime: now, + }) + middleShareID := insertPeerShare(t, providerRepo, &sqlite.PeerShare{ + Name: "middle-share", + NodeID: providerMiddleNodeID, + Token: "share-middle-token", + PortRangeStart: 44000, + PortRangeEnd: 44010, + IsActive: 1, + CreatedTime: now, + UpdatedTime: now, + }) + exitShareID := insertPeerShare(t, providerRepo, &sqlite.PeerShare{ + Name: "exit-share", + NodeID: providerExitNodeID, + Token: "share-exit-token", + PortRangeStart: 45000, + PortRangeEnd: 45010, + IsActive: 1, + CreatedTime: now, + UpdatedTime: now, + }) + + importRemoteNodeForContract(t, consumerRouter, consumerAdminToken, providerServer.URL, "share-entry-token") + importRemoteNodeForContract(t, consumerRouter, consumerAdminToken, providerServer.URL, "share-middle-token") + importRemoteNodeForContract(t, consumerRouter, consumerAdminToken, providerServer.URL, "share-exit-token") + + entryRemoteNodeID := queryRemoteNodeIDByToken(t, consumerRepo, "share-entry-token") + middleRemoteNodeID := queryRemoteNodeIDByToken(t, consumerRepo, "share-middle-token") + exitRemoteNodeID := queryRemoteNodeIDByToken(t, consumerRepo, "share-exit-token") + + stopMiddle := startMockNodeSession(t, providerServer.URL, "provider-middle-secret") + defer stopMiddle() + stopExit := startMockNodeSession(t, providerServer.URL, "provider-exit-secret") + defer stopExit() + + createTunnel := func(name string) int64 { + payload := map[string]interface{}{ + "name": name, + "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(payload) + if err != nil { + t.Fatalf("marshal create payload: %v", err) + } + req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/create", bytes.NewReader(body)) + req.Header.Set("Authorization", consumerAdminToken) + req.Header.Set("Content-Type", "application/json") + res := httptest.NewRecorder() + consumerRouter.ServeHTTP(res, req) + assertCode(t, res, 0) + + var tunnelID int64 + if err := consumerRepo.DB().QueryRow(`SELECT id FROM tunnel WHERE name = ? ORDER BY id DESC LIMIT 1`, name).Scan(&tunnelID); err != nil { + t.Fatalf("query tunnel id (%s): %v", name, err) + } + if tunnelID <= 0 { + t.Fatalf("invalid tunnel id for %s", name) + } + return tunnelID + } + + firstTunnelID := createTunnel("dual-panel-middle-exit-1") + + assertTunnelPortInRange(t, consumerRepo, firstTunnelID, 2, middleRemoteNodeID, 44000, 44010) + assertTunnelPortInRange(t, consumerRepo, firstTunnelID, 3, exitRemoteNodeID, 45000, 45010) + + assertCount(t, consumerRepo, `SELECT COUNT(1) FROM federation_tunnel_binding WHERE tunnel_id = ? AND status = 1`, firstTunnelID, 2) + 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) + + deleteBody, err := json.Marshal(map[string]interface{}{"id": firstTunnelID}) + if err != nil { + t.Fatalf("marshal delete payload: %v", err) + } + deleteReq := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/delete", bytes.NewReader(deleteBody)) + deleteReq.Header.Set("Authorization", consumerAdminToken) + deleteReq.Header.Set("Content-Type", "application/json") + deleteRes := httptest.NewRecorder() + consumerRouter.ServeHTTP(deleteRes, deleteReq) + assertCode(t, deleteRes, 0) + + assertCount(t, consumerRepo, `SELECT COUNT(1) FROM federation_tunnel_binding WHERE tunnel_id = ?`, firstTunnelID, 0) + assertCount(t, providerRepo, `SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ? AND status = 0`, middleShareID, 1) + assertCount(t, providerRepo, `SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ? AND status = 0`, exitShareID, 1) + + secondTunnelID := createTunnel("dual-panel-middle-exit-2") + assertTunnelPortInRange(t, consumerRepo, secondTunnelID, 2, middleRemoteNodeID, 44000, 44010) + assertTunnelPortInRange(t, consumerRepo, secondTunnelID, 3, exitRemoteNodeID, 45000, 45010) + + 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) +} + +func insertContractNode(t *testing.T, repo *sqlite.Repository, name, ip, portRange, secret string, status int) int64 { + t.Helper() + now := time.Now().UnixMilli() + 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, secret, ip, ip, "", portRange, "", "v1", 1, 1, 1, now, now, status, "[::]", "[::]", 0) + if err != nil { + t.Fatalf("insert node %s: %v", name, err) + } + id, err := res.LastInsertId() + if err != nil { + t.Fatalf("node id %s: %v", name, err) + } + return id +} + +func insertPeerShare(t *testing.T, repo *sqlite.Repository, share *sqlite.PeerShare) int64 { + t.Helper() + if share == nil { + t.Fatalf("share is nil") + } + if err := repo.CreatePeerShare(share); err != nil { + t.Fatalf("create peer share %s: %v", share.Name, err) + } + saved, err := repo.GetPeerShareByToken(share.Token) + if err != nil { + t.Fatalf("query peer share %s: %v", share.Name, err) + } + if saved == nil { + t.Fatalf("peer share %s not found after create", share.Name) + } + return saved.ID +} + +func importRemoteNodeForContract(t *testing.T, router http.Handler, adminToken, remoteURL, token string) { + t.Helper() + body, err := json.Marshal(map[string]string{ + "remoteUrl": remoteURL, + "token": token, + }) + if err != nil { + t.Fatalf("marshal import payload: %v", err) + } + req := httptest.NewRequest(http.MethodPost, "/api/v1/federation/node/import", bytes.NewReader(body)) + req.Header.Set("Authorization", adminToken) + req.Header.Set("Content-Type", "application/json") + res := httptest.NewRecorder() + router.ServeHTTP(res, req) + assertCode(t, res, 0) +} + +func queryRemoteNodeIDByToken(t *testing.T, repo *sqlite.Repository, token string) int64 { + t.Helper() + var id int64 + if err := repo.DB().QueryRow(`SELECT id FROM node WHERE is_remote = 1 AND remote_token = ? ORDER BY id DESC LIMIT 1`, token).Scan(&id); err != nil { + t.Fatalf("query remote node by token %s: %v", token, err) + } + if id <= 0 { + t.Fatalf("invalid remote node id for token %s", token) + } + return id +} + +func assertTunnelPortInRange(t *testing.T, repo *sqlite.Repository, tunnelID int64, chainType int, nodeID int64, minPort int, maxPort int) { + t.Helper() + var port int + err := repo.DB().QueryRow(` + SELECT port + FROM chain_tunnel + WHERE tunnel_id = ? AND chain_type = ? AND node_id = ? + LIMIT 1 + `, tunnelID, chainType, nodeID).Scan(&port) + if err != nil { + t.Fatalf("query tunnel=%d chainType=%d node=%d port: %v", tunnelID, chainType, nodeID, err) + } + if port < minPort || port > maxPort { + t.Fatalf("expected port in range [%d,%d], got %d", minPort, maxPort, port) + } +} + +func assertCount(t *testing.T, repo *sqlite.Repository, query string, arg interface{}, expected int) { + t.Helper() + var got int + if err := repo.DB().QueryRow(query, arg).Scan(&got); err != nil { + t.Fatalf("count query failed: %v", err) + } + if got != expected { + t.Fatalf("expected count %d, got %d (query: %s, arg: %v)", expected, got, query, arg) + } +} + +func startMockNodeSession(t *testing.T, baseURL string, nodeSecret string) func() { + t.Helper() + u, err := url.Parse(baseURL) + if err != nil { + t.Fatalf("parse provider url: %v", err) + } + if strings.EqualFold(u.Scheme, "https") { + u.Scheme = "wss" + } else { + u.Scheme = "ws" + } + u.Path = "/system-info" + q := u.Query() + q.Set("type", "1") + q.Set("secret", nodeSecret) + q.Set("version", "v1") + q.Set("http", "1") + q.Set("tls", "1") + q.Set("socks", "1") + u.RawQuery = q.Encode() + + conn, _, err := websocket.DefaultDialer.Dial(u.String(), nil) + if err != nil { + t.Fatalf("dial mock node websocket: %v", err) + } + + var wg sync.WaitGroup + wg.Add(1) + go func() { + defer wg.Done() + for { + _, raw, readErr := conn.ReadMessage() + if readErr != nil { + return + } + + plain := raw + var wrap struct { + Encrypted bool `json:"encrypted"` + Data string `json:"data"` + } + if err := json.Unmarshal(raw, &wrap); err == nil && wrap.Encrypted && strings.TrimSpace(wrap.Data) != "" { + crypto, cryptoErr := security.NewAESCrypto(nodeSecret) + if cryptoErr == nil { + if dec, decErr := crypto.Decrypt(wrap.Data); decErr == nil { + plain = []byte(dec) + } + } + } + + var cmd struct { + Type string `json:"type"` + RequestID string `json:"requestId"` + } + if err := json.Unmarshal(plain, &cmd); err != nil { + continue + } + if strings.TrimSpace(cmd.RequestID) == "" { + continue + } + + respType := fmt.Sprintf("%sResponse", cmd.Type) + respBytes, err := json.Marshal(map[string]interface{}{ + "type": respType, + "success": true, + "message": "OK", + "requestId": cmd.RequestID, + }) + if err != nil { + continue + } + _ = conn.WriteMessage(websocket.TextMessage, respBytes) + } + }() + + return func() { + _ = conn.Close() + wg.Wait() + } +}