diff --git a/docker-compose-v4.yml b/docker-compose-v4.yml index ea85469..3e4bbe0 100644 --- a/docker-compose-v4.yml +++ b/docker-compose-v4.yml @@ -87,4 +87,4 @@ networks: driver: bridge ipam: config: - - subnet: 172.20.0.0/16 + - subnet: 172.80.0.0/16 diff --git a/docker-compose-v6.yml b/docker-compose-v6.yml index cb7ae47..b77fc2f 100644 --- a/docker-compose-v6.yml +++ b/docker-compose-v6.yml @@ -88,5 +88,5 @@ networks: enable_ipv6: true ipam: config: - - subnet: 172.20.0.0/16 + - subnet: 172.80.0.0/16 - subnet: fd00:dead:beef::/48 diff --git a/go-backend/internal/http/client/federation.go b/go-backend/internal/http/client/federation.go index 128055a..09d6a64 100644 --- a/go-backend/internal/http/client/federation.go +++ b/go-backend/internal/http/client/federation.go @@ -77,6 +77,18 @@ type RuntimeDiagnoseRequest struct { Timeout int `json:"timeout"` } +type RuntimeNodeCommandRequest struct { + CommandType string `json:"commandType"` + Data interface{} `json:"data"` +} + +type RuntimeNodeCommandResponse struct { + Type string `json:"type"` + Success bool `json:"success"` + Message string `json:"message"` + Data map[string]interface{} `json:"data,omitempty"` +} + func NewFederationClient() *FederationClient { return &FederationClient{ client: &http.Client{ @@ -333,3 +345,42 @@ func (c *FederationClient) Diagnose(url, token, localDomain string, reqData Runt return res.Data, nil } + +func (c *FederationClient) Command(url, token, localDomain string, reqData RuntimeNodeCommandRequest) (*RuntimeNodeCommandResponse, error) { + url = strings.TrimSuffix(url, "/") + bodyBytes, _ := json.Marshal(reqData) + req, err := http.NewRequest("POST", url+"/api/v1/federation/runtime/command", 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 RuntimeNodeCommandResponse `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) + } + + 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 0e3d084..d61a86a 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -457,7 +457,17 @@ func (h *Handler) applyNodeProtocolChange(nodeID int64, httpVal, tlsVal, socksVa } func (h *Handler) sendNodeCommand(nodeID int64, commandType string, data interface{}, tolerateExists bool, tolerateNotFound bool) (ws.CommandResult, error) { - result, err := h.wsServer.SendCommand(nodeID, commandType, data, 12*time.Second) + var ( + result ws.CommandResult + err error + ) + + node, nodeErr := h.getNodeRecord(nodeID) + if nodeErr == nil && node != nil && node.IsRemote == 1 { + result, err = h.sendRemoteNodeCommand(node, commandType, data) + } else { + result, err = h.wsServer.SendCommand(nodeID, commandType, data, 12*time.Second) + } if err == nil { return result, nil } @@ -475,6 +485,44 @@ func (h *Handler) sendNodeCommand(nodeID int64, commandType string, data interfa return result, err } +func (h *Handler) sendRemoteNodeCommand(node *nodeRecord, commandType string, data interface{}) (ws.CommandResult, error) { + if node == nil { + return ws.CommandResult{}, errors.New("节点不存在") + } + remoteURL := strings.TrimSpace(node.RemoteURL) + remoteToken := strings.TrimSpace(node.RemoteToken) + if remoteURL == "" || remoteToken == "" { + return ws.CommandResult{}, errors.New("远程节点缺少共享配置") + } + + fc := client.NewFederationClient() + res, err := fc.Command(remoteURL, remoteToken, h.federationLocalDomain(), client.RuntimeNodeCommandRequest{ + CommandType: commandType, + Data: data, + }) + if err != nil { + return ws.CommandResult{}, err + } + if res == nil { + return ws.CommandResult{}, errors.New("远程节点未返回命令结果") + } + + result := ws.CommandResult{ + Type: res.Type, + Success: res.Success, + Message: res.Message, + Data: res.Data, + } + if !result.Success { + msg := strings.TrimSpace(result.Message) + if msg == "" { + msg = "命令执行失败" + } + return result, errors.New(msg) + } + return result, nil +} + func (h *Handler) diagnoseForwardRuntime(forward *forwardRecord) (map[string]interface{}, error) { if forward == nil { return nil, errForwardNotFound diff --git a/go-backend/internal/http/handler/federation.go b/go-backend/internal/http/handler/federation.go index 09abef2..831f64f 100644 --- a/go-backend/internal/http/handler/federation.go +++ b/go-backend/internal/http/handler/federation.go @@ -91,6 +91,11 @@ type federationRuntimeDiagnoseRequest struct { Timeout int `json:"timeout"` } +type federationRuntimeCommandRequest struct { + CommandType string `json:"commandType"` + Data interface{} `json:"data"` +} + type peerShareUsedPort struct { RuntimeID int64 `json:"runtimeId"` Port int `json:"port"` @@ -1199,6 +1204,51 @@ func (h *Handler) federationRuntimeDiagnose(w http.ResponseWriter, r *http.Reque response.WriteJSON(w, response.OK(res.Data)) } +func (h *Handler) federationRuntimeCommand(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 federationRuntimeCommandRequest + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("Invalid JSON")) + return + } + cmd := strings.TrimSpace(req.CommandType) + if cmd == "" { + response.WriteJSON(w, response.ErrDefault("commandType is required")) + return + } + if !isFederationRuntimeCommandAllowed(cmd) { + response.WriteJSON(w, response.ErrDefault("command not allowed")) + return + } + + res, err := h.sendNodeCommand(share.NodeID, cmd, req.Data, false, false) + if err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + response.WriteJSON(w, response.OK(res)) +} + +func isFederationRuntimeCommandAllowed(commandType string) bool { + switch strings.ToLower(strings.TrimSpace(commandType)) { + case "addservice", "updateservice", "deleteservice", "pauseservice", "resumeservice", "addchains", "deletechains", "addlimiters", "deletelimiters", "tcpping", "reload": + return true + default: + return false + } +} + func (h *Handler) pickPeerSharePort(share *sqlite.PeerShare, requestedPort int) (int, error) { if share == nil { return 0, fmt.Errorf("share not found") diff --git a/go-backend/internal/http/handler/federation_runtime_test.go b/go-backend/internal/http/handler/federation_runtime_test.go index b29f758..17c1520 100644 --- a/go-backend/internal/http/handler/federation_runtime_test.go +++ b/go-backend/internal/http/handler/federation_runtime_test.go @@ -152,6 +152,71 @@ func TestPrepareTunnelCreateStateRemoteAutoPortDefersToFederation(t *testing.T) } } +func TestPrepareTunnelCreateStateAllowsOfflineRemoteMiddleNode(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":2}`) + 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-local", 1, "32000-32010", 0) + remoteMiddleID := insertNode("middle-remote", 0, "33000-33010", 1) + outID := insertNode("out-local", 1, "34000-34010", 0) + + tx, err := repo.DB().Begin() + if err != nil { + t.Fatalf("begin tx: %v", err) + } + defer tx.Rollback() + + req := map[string]interface{}{ + "name": "remote-middle-offline-status", + "inNodeId": []interface{}{ + map[string]interface{}{"nodeId": float64(entryID), "protocol": "tls", "strategy": "round"}, + }, + "chainNodes": []interface{}{ + []interface{}{ + map[string]interface{}{"nodeId": float64(remoteMiddleID), "protocol": "tls", "strategy": "round", "port": float64(0)}, + }, + }, + "outNodeId": []interface{}{ + map[string]interface{}{"nodeId": float64(outID), "protocol": "tls", "strategy": "round", "port": float64(0)}, + }, + } + + state, err := h.prepareTunnelCreateState(tx, req, 2, 0) + if err != nil { + t.Fatalf("prepare state should allow offline remote middle node: %v", err) + } + if len(state.ChainHops) != 1 || len(state.ChainHops[0]) != 1 { + t.Fatalf("expected one middle hop node, got %+v", state.ChainHops) + } + if state.ChainHops[0][0].NodeID != remoteMiddleID { + t.Fatalf("expected remote middle node id %d, got %d", remoteMiddleID, state.ChainHops[0][0].NodeID) + } + if state.Nodes[remoteMiddleID] == nil || state.Nodes[remoteMiddleID].IsRemote != 1 { + t.Fatalf("expected remote middle node metadata in state") + } +} + func TestFederationRuntimeReservePortRejectsWhenShareFlowExceeded(t *testing.T) { repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db")) if err != nil { diff --git a/go-backend/internal/http/handler/handler.go b/go-backend/internal/http/handler/handler.go index 51c4093..03530d5 100644 --- a/go-backend/internal/http/handler/handler.go +++ b/go-backend/internal/http/handler/handler.go @@ -169,6 +169,7 @@ func (h *Handler) Register(mux *http.ServeMux) { 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/runtime/command", h.authPeer(h.federationRuntimeCommand)) mux.HandleFunc("/api/v1/federation/node/import", h.nodeImport) mux.HandleFunc("/api/v1/backup/export", h.backupExport) diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index cff0744..c2cc211 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -1713,10 +1713,19 @@ func (h *Handler) groupUserAssign(w http.ResponseWriter, r *http.Request) { return } defer func() { _ = tx.Rollback() }() + previousUserIDs, err := queryInt64ListTx(tx, `SELECT user_id FROM user_group_user WHERE user_group_id = ?`, req.GroupID) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } _, _ = tx.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, req.GroupID) for _, uid := range req.UserIDs { _, _ = tx.Exec(`INSERT INTO user_group_user(user_group_id, user_id, created_time) VALUES(?, ?, ?) ON CONFLICT DO NOTHING`, req.GroupID, uid, time.Now().UnixMilli()) } + if err := revokeGroupGrantsForRemovedUsersTx(tx, req.GroupID, previousUserIDs, req.UserIDs); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } if err := tx.Commit(); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return @@ -1748,10 +1757,35 @@ func (h *Handler) groupPermissionRemove(w http.ResponseWriter, r *http.Request) if id <= 0 { return } + tx, err := h.repo.DB().Begin() + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + defer func() { _ = tx.Rollback() }() + var ug, tg int64 - _ = h.repo.DB().QueryRow(`SELECT user_group_id, tunnel_group_id FROM group_permission WHERE id = ?`, id).Scan(&ug, &tg) - _, _ = h.repo.DB().Exec(`DELETE FROM group_permission WHERE id = ?`, id) - _, _ = h.repo.DB().Exec(`DELETE FROM group_permission_grant WHERE user_group_id = ? AND tunnel_group_id = ?`, ug, tg) + err = tx.QueryRow(`SELECT user_group_id, tunnel_group_id FROM group_permission WHERE id = ?`, id).Scan(&ug, &tg) + if err != nil && err != sql.ErrNoRows { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + + if _, err := tx.Exec(`DELETE FROM group_permission WHERE id = ?`, id); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if err == nil { + if err := revokeGroupPermissionPairTx(tx, ug, tg); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + } + + if err := tx.Commit(); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } response.WriteJSON(w, response.OKEmpty()) } @@ -1908,6 +1942,144 @@ func queryInt64List(db *store.DB, q string, args ...interface{}) ([]int64, error return out, rows.Err() } +func queryInt64ListTx(tx *store.Tx, q string, args ...interface{}) ([]int64, error) { + rows, err := tx.Query(q, args...) + if err != nil { + return nil, err + } + defer rows.Close() + out := make([]int64, 0) + for rows.Next() { + var v int64 + if err := rows.Scan(&v); err != nil { + return nil, err + } + out = append(out, v) + } + return out, rows.Err() +} + +func revokeGroupGrantsForRemovedUsersTx(tx *store.Tx, userGroupID int64, previousUserIDs, currentUserIDs []int64) error { + currentSet := make(map[int64]struct{}, len(currentUserIDs)) + for _, uid := range currentUserIDs { + if uid > 0 { + currentSet[uid] = struct{}{} + } + } + + removedUserIDs := make([]int64, 0) + for _, uid := range previousUserIDs { + if uid <= 0 { + continue + } + if _, ok := currentSet[uid]; !ok { + removedUserIDs = append(removedUserIDs, uid) + } + } + if len(removedUserIDs) == 0 { + return nil + } + + for _, userID := range removedUserIDs { + rows, err := tx.Query(` + SELECT g.user_tunnel_id, g.created_by_group + FROM group_permission_grant g + JOIN user_tunnel ut ON ut.id = g.user_tunnel_id + WHERE g.user_group_id = ? AND ut.user_id = ? + `, userGroupID, userID) + if err != nil { + return err + } + + groupCreatedTunnelIDs := make(map[int64]struct{}) + for rows.Next() { + var userTunnelID int64 + var createdByGroup int + if err := rows.Scan(&userTunnelID, &createdByGroup); err != nil { + rows.Close() + return err + } + if createdByGroup == 1 && userTunnelID > 0 { + groupCreatedTunnelIDs[userTunnelID] = struct{}{} + } + } + if err := rows.Err(); err != nil { + rows.Close() + return err + } + rows.Close() + + if _, err := tx.Exec(` + DELETE FROM group_permission_grant + WHERE user_group_id = ? + AND user_tunnel_id IN (SELECT id FROM user_tunnel WHERE user_id = ?) + `, userGroupID, userID); err != nil { + return err + } + + for userTunnelID := range groupCreatedTunnelIDs { + var remaining int + if err := tx.QueryRow(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Scan(&remaining); err != nil { + return err + } + if remaining == 0 { + if _, err := tx.Exec(`DELETE FROM user_tunnel WHERE id = ?`, userTunnelID); err != nil { + return err + } + } + } + } + + return nil +} + +func revokeGroupPermissionPairTx(tx *store.Tx, userGroupID, tunnelGroupID int64) error { + rows, err := tx.Query(` + SELECT user_tunnel_id, created_by_group + FROM group_permission_grant + WHERE user_group_id = ? AND tunnel_group_id = ? + `, userGroupID, tunnelGroupID) + if err != nil { + return err + } + + groupCreatedTunnelIDs := make(map[int64]struct{}) + for rows.Next() { + var userTunnelID int64 + var createdByGroup int + if err := rows.Scan(&userTunnelID, &createdByGroup); err != nil { + rows.Close() + return err + } + if createdByGroup == 1 && userTunnelID > 0 { + groupCreatedTunnelIDs[userTunnelID] = struct{}{} + } + } + if err := rows.Err(); err != nil { + rows.Close() + return err + } + rows.Close() + + if _, err := tx.Exec(`DELETE FROM group_permission_grant WHERE user_group_id = ? AND tunnel_group_id = ?`, userGroupID, tunnelGroupID); err != nil { + return err + } + + for userTunnelID := range groupCreatedTunnelIDs { + var remaining int + if err := tx.QueryRow(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Scan(&remaining); err != nil { + return err + } + if remaining == 0 { + if _, err := tx.Exec(`DELETE FROM user_tunnel WHERE id = ?`, userTunnelID); err != nil { + return err + } + } + } + + return nil +} + func queryPairs(db *store.DB, q string, args ...interface{}) ([][2]int64, error) { rows, err := db.Query(q, args...) if err != nil { @@ -2061,7 +2233,7 @@ func (h *Handler) prepareTunnelCreateState(tx *store.Tx, req map[string]interfac } return nil, err } - if node.Status != 1 { + if node.IsRemote != 1 && node.Status != 1 { return nil, errors.New("部分节点不在线") } state.Nodes[nodeID] = node diff --git a/go-backend/internal/http/middleware/auth.go b/go-backend/internal/http/middleware/auth.go index bf3aebf..7e60afc 100644 --- a/go-backend/internal/http/middleware/auth.go +++ b/go-backend/internal/http/middleware/auth.go @@ -93,6 +93,8 @@ func shouldSkip(path string) bool { return true case path == "/api/v1/federation/runtime/diagnose": return true + case path == "/api/v1/federation/runtime/command": + return true default: return false } 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 b3bece9..0959f65 100644 --- a/go-backend/tests/contract/federation_dual_panel_contract_test.go +++ b/go-backend/tests/contract/federation_dual_panel_contract_test.go @@ -78,6 +78,8 @@ func TestFederationDualPanelMiddleExitAutoPortContract(t *testing.T) { middleRemoteNodeID := queryRemoteNodeIDByToken(t, consumerRepo, "share-middle-token") exitRemoteNodeID := queryRemoteNodeIDByToken(t, consumerRepo, "share-exit-token") + stopEntry := startMockNodeSession(t, providerServer.URL, "provider-entry-secret") + defer stopEntry() stopMiddle := startMockNodeSession(t, providerServer.URL, "provider-middle-secret") defer stopMiddle() stopExit := startMockNodeSession(t, providerServer.URL, "provider-exit-secret") @@ -148,6 +150,23 @@ func TestFederationDualPanelMiddleExitAutoPortContract(t *testing.T) { assertTunnelPortInRange(t, consumerRepo, secondTunnelID, 2, middleRemoteNodeID, 44000, 44010) assertTunnelPortInRange(t, consumerRepo, secondTunnelID, 3, exitRemoteNodeID, 45000, 45010) + forwardPayload := map[string]interface{}{ + "name": "dual-panel-remote-entry-forward", + "tunnelId": secondTunnelID, + "remoteAddr": "1.1.1.1:443", + "strategy": "fifo", + } + forwardBody, err := json.Marshal(forwardPayload) + if err != nil { + t.Fatalf("marshal forward payload: %v", err) + } + forwardReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(forwardBody)) + forwardReq.Header.Set("Authorization", consumerAdminToken) + forwardReq.Header.Set("Content-Type", "application/json") + forwardRes := httptest.NewRecorder() + consumerRouter.ServeHTTP(forwardRes, forwardReq) + assertCode(t, forwardRes, 0) + 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) diff --git a/go-backend/tests/contract/group_permission_contract_test.go b/go-backend/tests/contract/group_permission_contract_test.go new file mode 100644 index 0000000..530acdd --- /dev/null +++ b/go-backend/tests/contract/group_permission_contract_test.go @@ -0,0 +1,219 @@ +package contract_test + +import ( + "bytes" + "net/http" + "net/http/httptest" + "testing" + "time" + + "go-backend/internal/auth" +) + +func TestGroupUserUnbindRevokesInheritedTunnelPermission(t *testing.T) { + secret := "contract-jwt-secret" + router, repo := setupContractRouter(t, secret) + now := time.Now().UnixMilli() + + if _, err := repo.DB().Exec(` + INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) + VALUES(200, 'group_user_contract', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1) + `, now, now); err != nil { + t.Fatalf("insert test user: %v", err) + } + + tunnelRes, err := repo.DB().Exec(` + INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) + VALUES('group-contract-tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0) + `, now, now) + if err != nil { + t.Fatalf("insert tunnel: %v", err) + } + tunnelID, err := tunnelRes.LastInsertId() + if err != nil { + t.Fatalf("read tunnel id: %v", err) + } + + ugRes, err := repo.DB().Exec(`INSERT INTO user_group(name, created_time, updated_time, status) VALUES('ug-contract', ?, ?, 1)`, now, now) + if err != nil { + t.Fatalf("insert user_group: %v", err) + } + userGroupID, err := ugRes.LastInsertId() + if err != nil { + t.Fatalf("read user_group id: %v", err) + } + + tgRes, err := repo.DB().Exec(`INSERT INTO tunnel_group(name, created_time, updated_time, status) VALUES('tg-contract', ?, ?, 1)`, now, now) + if err != nil { + t.Fatalf("insert tunnel_group: %v", err) + } + tunnelGroupID, err := tgRes.LastInsertId() + if err != nil { + t.Fatalf("read tunnel_group id: %v", err) + } + + if _, err := repo.DB().Exec(`INSERT INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time) VALUES(?, ?, ?)`, tunnelGroupID, tunnelID, now); err != nil { + t.Fatalf("insert tunnel_group_tunnel: %v", err) + } + if _, err := repo.DB().Exec(`INSERT INTO group_permission(user_group_id, tunnel_group_id, created_time) VALUES(?, ?, ?)`, userGroupID, tunnelGroupID, now); err != nil { + t.Fatalf("insert group_permission: %v", err) + } + + adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret) + if err != nil { + t.Fatalf("generate admin token: %v", err) + } + + bindReq := httptest.NewRequest(http.MethodPost, "/api/v1/group/user/assign", bytes.NewBufferString(`{"groupId":`+jsonNumber(userGroupID)+`,"userIds":[200]}`)) + bindReq.Header.Set("Authorization", adminToken) + bindRes := httptest.NewRecorder() + router.ServeHTTP(bindRes, bindReq) + assertCode(t, bindRes, 0) + + var userTunnelID int64 + if err := repo.DB().QueryRow(`SELECT id FROM user_tunnel WHERE user_id = 200 AND tunnel_id = ?`, tunnelID).Scan(&userTunnelID); err != nil { + t.Fatalf("query user_tunnel after bind: %v", err) + } + + var grantCount int + if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Scan(&grantCount); err != nil { + t.Fatalf("query group_permission_grant after bind: %v", err) + } + if grantCount == 0 { + t.Fatalf("expected non-zero grants after bind") + } + + unbindReq := httptest.NewRequest(http.MethodPost, "/api/v1/group/user/assign", bytes.NewBufferString(`{"groupId":`+jsonNumber(userGroupID)+`,"userIds":[]}`)) + unbindReq.Header.Set("Authorization", adminToken) + unbindRes := httptest.NewRecorder() + router.ServeHTTP(unbindRes, unbindReq) + assertCode(t, unbindRes, 0) + + if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Scan(&grantCount); err != nil { + t.Fatalf("query group_permission_grant after unbind: %v", err) + } + if grantCount != 0 { + t.Fatalf("expected grants revoked after unbind, got %d", grantCount) + } + + var userTunnelCount int + if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM user_tunnel WHERE id = ?`, userTunnelID).Scan(&userTunnelCount); err != nil { + t.Fatalf("query user_tunnel after unbind: %v", err) + } + if userTunnelCount != 0 { + t.Fatalf("expected user_tunnel revoked after unbind, got %d", userTunnelCount) + } +} + +func TestGroupPermissionRemoveRevokesInheritedTunnelPermission(t *testing.T) { + secret := "contract-jwt-secret" + router, repo := setupContractRouter(t, secret) + now := time.Now().UnixMilli() + + if _, err := repo.DB().Exec(` + INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) + VALUES(201, 'group_user_permission_remove', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1) + `, now, now); err != nil { + t.Fatalf("insert test user: %v", err) + } + + tunnelRes, err := repo.DB().Exec(` + INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) + VALUES('group-remove-tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0) + `, now, now) + if err != nil { + t.Fatalf("insert tunnel: %v", err) + } + tunnelID, err := tunnelRes.LastInsertId() + if err != nil { + t.Fatalf("read tunnel id: %v", err) + } + + ugRes, err := repo.DB().Exec(`INSERT INTO user_group(name, created_time, updated_time, status) VALUES('ug-remove-contract', ?, ?, 1)`, now, now) + if err != nil { + t.Fatalf("insert user_group: %v", err) + } + userGroupID, err := ugRes.LastInsertId() + if err != nil { + t.Fatalf("read user_group id: %v", err) + } + + tgRes, err := repo.DB().Exec(`INSERT INTO tunnel_group(name, created_time, updated_time, status) VALUES('tg-remove-contract', ?, ?, 1)`, now, now) + if err != nil { + t.Fatalf("insert tunnel_group: %v", err) + } + tunnelGroupID, err := tgRes.LastInsertId() + if err != nil { + t.Fatalf("read tunnel_group id: %v", err) + } + + adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret) + if err != nil { + t.Fatalf("generate admin token: %v", err) + } + + assignTunnelReq := httptest.NewRequest(http.MethodPost, "/api/v1/group/tunnel/assign", bytes.NewBufferString(`{"groupId":`+jsonNumber(tunnelGroupID)+`,"tunnelIds":[`+jsonNumber(tunnelID)+`]}`)) + assignTunnelReq.Header.Set("Authorization", adminToken) + assignTunnelRes := httptest.NewRecorder() + router.ServeHTTP(assignTunnelRes, assignTunnelReq) + assertCode(t, assignTunnelRes, 0) + + assignUserReq := httptest.NewRequest(http.MethodPost, "/api/v1/group/user/assign", bytes.NewBufferString(`{"groupId":`+jsonNumber(userGroupID)+`,"userIds":[201]}`)) + assignUserReq.Header.Set("Authorization", adminToken) + assignUserRes := httptest.NewRecorder() + router.ServeHTTP(assignUserRes, assignUserReq) + assertCode(t, assignUserRes, 0) + + assignPermissionReq := httptest.NewRequest(http.MethodPost, "/api/v1/group/permission/assign", bytes.NewBufferString(`{"userGroupId":`+jsonNumber(userGroupID)+`,"tunnelGroupId":`+jsonNumber(tunnelGroupID)+`}`)) + assignPermissionReq.Header.Set("Authorization", adminToken) + assignPermissionRes := httptest.NewRecorder() + router.ServeHTTP(assignPermissionRes, assignPermissionReq) + assertCode(t, assignPermissionRes, 0) + + var permissionID int64 + if err := repo.DB().QueryRow(`SELECT id FROM group_permission WHERE user_group_id = ? AND tunnel_group_id = ?`, userGroupID, tunnelGroupID).Scan(&permissionID); err != nil { + t.Fatalf("query group_permission id: %v", err) + } + + var userTunnelID int64 + if err := repo.DB().QueryRow(`SELECT id FROM user_tunnel WHERE user_id = 201 AND tunnel_id = ?`, tunnelID).Scan(&userTunnelID); err != nil { + t.Fatalf("query user_tunnel after assign: %v", err) + } + + var grantCount int + if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Scan(&grantCount); err != nil { + t.Fatalf("query group_permission_grant after assign: %v", err) + } + if grantCount == 0 { + t.Fatalf("expected non-zero grants after permission assign") + } + + removeReq := httptest.NewRequest(http.MethodPost, "/api/v1/group/permission/remove", bytes.NewBufferString(`{"id":`+jsonNumber(permissionID)+`}`)) + removeReq.Header.Set("Authorization", adminToken) + removeRes := httptest.NewRecorder() + router.ServeHTTP(removeRes, removeReq) + assertCode(t, removeRes, 0) + + var permissionCount int + if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM group_permission WHERE id = ?`, permissionID).Scan(&permissionCount); err != nil { + t.Fatalf("query group_permission after remove: %v", err) + } + if permissionCount != 0 { + t.Fatalf("expected group_permission removed, got %d", permissionCount) + } + + if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Scan(&grantCount); err != nil { + t.Fatalf("query group_permission_grant after remove: %v", err) + } + if grantCount != 0 { + t.Fatalf("expected grants removed after permission remove, got %d", grantCount) + } + + var userTunnelCount int + if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM user_tunnel WHERE id = ?`, userTunnelID).Scan(&userTunnelCount); err != nil { + t.Fatalf("query user_tunnel after permission remove: %v", err) + } + if userTunnelCount != 0 { + t.Fatalf("expected user_tunnel revoked after permission remove, got %d", userTunnelCount) + } +}