diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index b2e124d..89f82e0 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -200,6 +200,26 @@ func (h *Handler) listForwardPorts(forwardID int64) ([]forwardPortRecord, error) return result, nil } +func (h *Handler) isTunnelSelectedTLSProtocol(tunnelID int64) (bool, error) { + row := h.repo.DB().QueryRow(` + SELECT protocol + FROM chain_tunnel + WHERE tunnel_id = ? AND chain_type = '3' + ORDER BY id ASC + LIMIT 1 + `, tunnelID) + + var protocol sql.NullString + if err := row.Scan(&protocol); err != nil { + if errors.Is(err, sql.ErrNoRows) { + return false, nil + } + return false, err + } + + return isTLSTunnelProtocol(protocol.String), nil +} + func (h *Handler) getNodeRecord(nodeID int64) (*nodeRecord, error) { row := h.repo.DB().QueryRow(` SELECT id, name, server_ip, server_ip_v4, server_ip_v6, status, port, tcp_listen_addr, udp_listen_addr, interface_name, is_remote, remote_url, remote_token, remote_config @@ -346,6 +366,10 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all return err } serviceBase := buildForwardServiceBase(forward.ID, forward.UserID, userTunnelID) + tunnelTLSProtocol, err := h.isTunnelSelectedTLSProtocol(forward.TunnelID) + if err != nil { + return err + } for _, fp := range ports { if limiterID != nil && speed != nil { @@ -356,7 +380,7 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all if err != nil { return err } - services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, limiterID) + services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, limiterID, tunnelTLSProtocol) _, err = h.sendNodeCommand(node.ID, method, services, true, false) if err != nil && allowFallbackAdd && method == "UpdateService" { _, err = h.sendNodeCommand(node.ID, "AddService", services, true, false) @@ -1095,7 +1119,7 @@ func isNotFoundError(err error) bool { return strings.Contains(msg, "not found") || strings.Contains(msg, "不存在") } -func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, limiterID *int64) []map[string]interface{} { +func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, limiterID *int64, tunnelTLSProtocol bool) []map[string]interface{} { protocols := []string{"tcp", "udp"} services := make([]map[string]interface{}, 0, 2) targets := splitRemoteTargets(forward.RemoteAddr) @@ -1128,7 +1152,11 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel }, } if protocol == "udp" { - service["listener"].(map[string]interface{})["metadata"] = map[string]interface{}{"keepAlive": true} + listenerMetadata := map[string]interface{}{"keepAlive": true} + if tunnelTLSProtocol { + listenerMetadata["ttl"] = "10s" + } + service["listener"].(map[string]interface{})["metadata"] = listenerMetadata } if tunnel != nil && tunnel.Type == 2 { service["handler"].(map[string]interface{})["chain"] = fmt.Sprintf("chains_%d", forward.TunnelID) diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index db0f316..a1f4f80 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -2613,9 +2613,7 @@ func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64 } for _, inNode := range state.InNodes { - if node := state.Nodes[inNode.NodeID]; node != nil && node.IsRemote == 1 { - continue - } + node := state.Nodes[inNode.NodeID] targets := state.OutNodes if len(state.ChainHops) > 0 { targets = state.ChainHops[0] @@ -2625,6 +2623,9 @@ func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64 return createdChains, createdServices, err } if _, err := h.sendNodeCommand(inNode.NodeID, "AddChains", chainData, true, false); err != nil { + if node != nil && node.IsRemote == 1 && shouldDeferTunnelRuntimeApplyError(err) { + continue + } return createdChains, createdServices, fmt.Errorf("入口节点 %s 下发转发链失败: %w", nodeDisplayName(state.Nodes[inNode.NodeID]), err) } createdChains = append(createdChains, inNode.NodeID) diff --git a/go-backend/internal/store/sqlite/repository.go b/go-backend/internal/store/sqlite/repository.go index ca69abc..9d196c3 100644 --- a/go-backend/internal/store/sqlite/repository.go +++ b/go-backend/internal/store/sqlite/repository.go @@ -2330,13 +2330,25 @@ func (r *Repository) exportTunnels() ([]TunnelBackup, error) { var tunnels []TunnelBackup for rows.Next() { var t TunnelBackup + var protocol sql.NullString + var updatedTime sql.NullInt64 var inIP sql.NullString - if err := rows.Scan(&t.ID, &t.Name, &t.TrafficRatio, &t.Type, &t.Protocol, &t.Flow, &t.CreatedTime, &t.UpdatedTime, &t.Status, &inIP, &t.Inx); err != nil { + var inx sql.NullInt64 + if err := rows.Scan(&t.ID, &t.Name, &t.TrafficRatio, &t.Type, &protocol, &t.Flow, &t.CreatedTime, &updatedTime, &t.Status, &inIP, &inx); err != nil { return nil, err } + if protocol.Valid { + t.Protocol = protocol.String + } + if updatedTime.Valid { + t.UpdatedTime = updatedTime.Int64 + } if inIP.Valid { t.InIP = inIP.String } + if inx.Valid { + t.Inx = int(inx.Int64) + } // Export chain tunnels chainTunnels, err := r.exportChainTunnels(t.ID) if err != nil { @@ -2362,12 +2374,23 @@ func (r *Repository) exportChainTunnels(tunnelID int64) ([]ChainTunnelBackup, er for rows.Next() { var ct ChainTunnelBackup var port sql.NullInt64 - if err := rows.Scan(&ct.ID, &ct.TunnelID, &ct.ChainType, &ct.NodeID, &port, &ct.Strategy, &ct.Inx, &ct.Protocol); err != nil { + var strategy, protocol sql.NullString + var inx sql.NullInt64 + if err := rows.Scan(&ct.ID, &ct.TunnelID, &ct.ChainType, &ct.NodeID, &port, &strategy, &inx, &protocol); err != nil { return nil, err } if port.Valid { ct.Port = int(port.Int64) } + if strategy.Valid { + ct.Strategy = strategy.String + } + if inx.Valid { + ct.Inx = int(inx.Int64) + } + if protocol.Valid { + ct.Protocol = protocol.String + } chainTunnels = append(chainTunnels, ct) } return chainTunnels, rows.Err() @@ -2386,9 +2409,21 @@ func (r *Repository) exportForwards() ([]ForwardBackup, error) { var forwards []ForwardBackup for rows.Next() { var f ForwardBackup - if err := rows.Scan(&f.ID, &f.UserID, &f.UserName, &f.Name, &f.TunnelID, &f.RemoteAddr, &f.Strategy, &f.InFlow, &f.OutFlow, &f.CreatedTime, &f.UpdatedTime, &f.Status, &f.Inx); err != nil { + var strategy sql.NullString + var updatedTime sql.NullInt64 + var inx sql.NullInt64 + if err := rows.Scan(&f.ID, &f.UserID, &f.UserName, &f.Name, &f.TunnelID, &f.RemoteAddr, &strategy, &f.InFlow, &f.OutFlow, &f.CreatedTime, &updatedTime, &f.Status, &inx); err != nil { return nil, err } + if strategy.Valid { + f.Strategy = strategy.String + } + if updatedTime.Valid { + f.UpdatedTime = updatedTime.Int64 + } + if inx.Valid { + f.Inx = int(inx.Int64) + } forwards = append(forwards, f) } return forwards, rows.Err() 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 0959f65..784ad29 100644 --- a/go-backend/tests/contract/federation_dual_panel_contract_test.go +++ b/go-backend/tests/contract/federation_dual_panel_contract_test.go @@ -316,6 +316,136 @@ func TestFederationDualPanelRemoteDiagnosisContract(t *testing.T) { } } +func TestFederationDualPanelRemoteEntryRuntimeContract(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-rt", "198.51.100.21", "43020-43030", "provider-entry-rt-secret", 1) + providerMiddleNodeID := insertContractNode(t, providerRepo, "provider-middle-rt", "198.51.100.22", "44020-44030", "provider-middle-rt-secret", 1) + providerExitNodeID := insertContractNode(t, providerRepo, "provider-exit-rt", "198.51.100.23", "45020-45030", "provider-exit-rt-secret", 1) + + insertPeerShare(t, providerRepo, &sqlite.PeerShare{ + Name: "entry-share-rt", + NodeID: providerEntryNodeID, + Token: "share-entry-rt-token", + PortRangeStart: 43020, + PortRangeEnd: 43030, + IsActive: 1, + CreatedTime: now, + UpdatedTime: now, + }) + insertPeerShare(t, providerRepo, &sqlite.PeerShare{ + Name: "middle-share-rt", + NodeID: providerMiddleNodeID, + Token: "share-middle-rt-token", + PortRangeStart: 44020, + PortRangeEnd: 44030, + IsActive: 1, + CreatedTime: now, + UpdatedTime: now, + }) + insertPeerShare(t, providerRepo, &sqlite.PeerShare{ + Name: "exit-share-rt", + NodeID: providerExitNodeID, + Token: "share-exit-rt-token", + PortRangeStart: 45020, + PortRangeEnd: 45030, + IsActive: 1, + CreatedTime: now, + UpdatedTime: now, + }) + + importRemoteNodeForContract(t, consumerRouter, consumerAdminToken, providerServer.URL, "share-entry-rt-token") + importRemoteNodeForContract(t, consumerRouter, consumerAdminToken, providerServer.URL, "share-middle-rt-token") + importRemoteNodeForContract(t, consumerRouter, consumerAdminToken, providerServer.URL, "share-exit-rt-token") + + entryRemoteNodeID := queryRemoteNodeIDByToken(t, consumerRepo, "share-entry-rt-token") + middleRemoteNodeID := queryRemoteNodeIDByToken(t, consumerRepo, "share-middle-rt-token") + exitRemoteNodeID := queryRemoteNodeIDByToken(t, consumerRepo, "share-exit-rt-token") + + var commandMu sync.Mutex + entryCommands := make([]string, 0, 8) + stopEntry := startMockNodeSessionWithHook(t, providerServer.URL, "provider-entry-rt-secret", func(cmdType string) { + commandMu.Lock() + entryCommands = append(entryCommands, cmdType) + commandMu.Unlock() + }) + defer stopEntry() + stopMiddle := startMockNodeSession(t, providerServer.URL, "provider-middle-rt-secret") + defer stopMiddle() + stopExit := startMockNodeSession(t, providerServer.URL, "provider-exit-rt-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 + } + + createTunnel("dual-panel-remote-entry-online") + + commandMu.Lock() + seenAddChains := false + seenCommands := append([]string(nil), entryCommands...) + for _, cmdType := range entryCommands { + if strings.EqualFold(strings.TrimSpace(cmdType), "AddChains") { + seenAddChains = true + break + } + } + commandMu.Unlock() + if !seenAddChains { + t.Fatalf("expected entry remote node to receive AddChains, commands=%v", seenCommands) + } + + stopEntry() + waitNodeStatus(t, providerRepo, providerEntryNodeID, 0) + + createTunnel("dual-panel-remote-entry-offline") +} + func insertContractNode(t *testing.T, repo *sqlite.Repository, name, ip, portRange, secret string, status int) int64 { t.Helper() now := time.Now().UnixMilli() @@ -409,6 +539,10 @@ func assertCount(t *testing.T, repo *sqlite.Repository, query string, arg interf } func startMockNodeSession(t *testing.T, baseURL string, nodeSecret string) func() { + return startMockNodeSessionWithHook(t, baseURL, nodeSecret, nil) +} + +func startMockNodeSessionWithHook(t *testing.T, baseURL string, nodeSecret string, onCommand func(cmdType string)) func() { t.Helper() u, err := url.Parse(baseURL) if err != nil { @@ -468,6 +602,9 @@ func startMockNodeSession(t *testing.T, baseURL string, nodeSecret string) func( if strings.TrimSpace(cmd.RequestID) == "" { continue } + if onCommand != nil { + onCommand(strings.TrimSpace(cmd.Type)) + } respType := fmt.Sprintf("%sResponse", cmd.Type) respPayload := map[string]interface{}{ @@ -492,9 +629,27 @@ func startMockNodeSession(t *testing.T, baseURL string, nodeSecret string) func( } }() + var stopOnce sync.Once return func() { - _ = conn.Close() - wg.Wait() + stopOnce.Do(func() { + _ = conn.Close() + wg.Wait() + }) + } +} + +func waitNodeStatus(t *testing.T, repo *sqlite.Repository, nodeID int64, expectedStatus int) { + t.Helper() + deadline := time.Now().Add(2 * time.Second) + for { + var status int + if err := repo.DB().QueryRow(`SELECT status FROM node WHERE id = ?`, nodeID).Scan(&status); err == nil && status == expectedStatus { + return + } + if time.Now().After(deadline) { + t.Fatalf("node %d status did not reach %d before timeout", nodeID, expectedStatus) + } + time.Sleep(20 * time.Millisecond) } } diff --git a/go-backend/tests/contract/migration_contract_test.go b/go-backend/tests/contract/migration_contract_test.go index 4bc80a3..7ff2bd9 100644 --- a/go-backend/tests/contract/migration_contract_test.go +++ b/go-backend/tests/contract/migration_contract_test.go @@ -310,6 +310,80 @@ func TestBackupExportImportRestoreContracts(t *testing.T) { t.Fatalf("expected restored config value v3, got %+v", cfg) } }) + + t.Run("backup export tolerates nullable legacy tunnel chain fields", func(t *testing.T) { + now := time.Now().UnixMilli() + res, err := repo.DB().Exec(` + INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, "legacy-null-chain", 1.0, 1, "tls", 1000, now, now, 1, nil, 1) + if err != nil { + t.Fatalf("seed tunnel for nullable chain export: %v", err) + } + tunnelID, err := res.LastInsertId() + if err != nil { + t.Fatalf("read tunnel id for nullable chain export: %v", err) + } + + if _, err := repo.DB().Exec(` + INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) + VALUES(?, ?, ?, ?, ?, ?, ?) + `, tunnelID, "1", 1, nil, nil, nil, nil); err != nil { + t.Fatalf("seed nullable chain_tunnel row: %v", err) + } + + req := httptest.NewRequest(http.MethodPost, "/api/v1/backup/export", bytes.NewBufferString(`{"types":["tunnels"]}`)) + req.Header.Set("Authorization", adminToken) + req.Header.Set("Content-Type", "application/json") + resp := httptest.NewRecorder() + router.ServeHTTP(resp, req) + + if resp.Code != http.StatusOK { + t.Fatalf("expected status 200, got %d", resp.Code) + } + + var payload struct { + Version string `json:"version"` + Tunnels []struct { + ID int64 `json:"id"` + ChainTunnels []struct { + Inx int `json:"inx"` + Strategy string `json:"strategy"` + Protocol string `json:"protocol"` + } `json:"chainTunnels"` + } `json:"tunnels"` + } + if err := json.NewDecoder(resp.Body).Decode(&payload); err != nil { + t.Fatalf("decode tunnels backup payload: %v", err) + } + if strings.TrimSpace(payload.Version) == "" { + t.Fatalf("expected backup payload version, got empty") + } + + found := false + for _, tunnel := range payload.Tunnels { + if tunnel.ID != tunnelID { + continue + } + if len(tunnel.ChainTunnels) != 1 { + t.Fatalf("expected one chain tunnel for seeded tunnel %d, got %d", tunnelID, len(tunnel.ChainTunnels)) + } + if tunnel.ChainTunnels[0].Inx != 0 { + t.Fatalf("expected nullable chain inx to export as 0, got %d", tunnel.ChainTunnels[0].Inx) + } + if tunnel.ChainTunnels[0].Strategy != "" { + t.Fatalf("expected nullable chain strategy to export as empty string, got %q", tunnel.ChainTunnels[0].Strategy) + } + if tunnel.ChainTunnels[0].Protocol != "" { + t.Fatalf("expected nullable chain protocol to export as empty string, got %q", tunnel.ChainTunnels[0].Protocol) + } + found = true + break + } + if !found { + t.Fatalf("expected seeded tunnel %d in backup export", tunnelID) + } + }) } type backupExportPayload struct { diff --git a/vite-frontend/src/styles/globals.css b/vite-frontend/src/styles/globals.css index dc22ac2..9526312 100644 --- a/vite-frontend/src/styles/globals.css +++ b/vite-frontend/src/styles/globals.css @@ -46,6 +46,28 @@ html, body { --safe-area-bottom: env(safe-area-inset-bottom, 0px); } +[data-slot="input-wrapper"] { + box-shadow: none; +} + +[data-slot="input-wrapper"]:focus-within:not([data-invalid="true"]), +[data-slot="input-wrapper"][data-focus="true"]:not([data-invalid="true"]), +[data-slot="input-wrapper"][data-focused="true"]:not([data-invalid="true"]), +button[data-slot="trigger"][data-focus="true"]:not([data-invalid="true"]), +button[data-slot="trigger"][data-open="true"]:not([data-invalid="true"]) { + border-color: var(--heroui-default-200, #e5e7eb) !important; + outline: none !important; + outline-offset: 0 !important; + box-shadow: none !important; +} + +[data-slot="input-wrapper"] input:focus, +[data-slot="input-wrapper"] textarea:focus { + outline: none !important; + box-shadow: none !important; + border-color: transparent !important; +} + .safe-top { padding-top: var(--safe-area-top); } @@ -85,4 +107,4 @@ html, body { } } -@config "../../tailwind.config.js" \ No newline at end of file +@config "../../tailwind.config.js"