mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 07:36:38 +08:00
Compare commits
8 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| f1cad30f44 | |||
| 2e05df288b | |||
| 3e5bb8fc0b | |||
| d1e3c59537 | |||
| 8b8ebb6092 | |||
| 0195a2a01b | |||
| ad9b336fb9 | |||
| 30d9552207 |
@@ -44,10 +44,8 @@ func (h *Handler) processFlowItem(nodeID int64, item flowItem) {
|
||||
if ok {
|
||||
inFlow, outFlow := h.scaleFlowByTunnel(forwardID, item.D, item.U)
|
||||
_ = h.repo.AddFlow(forwardID, userID, userTunnelID, inFlow, outFlow)
|
||||
if forward, err := h.getForwardRecord(forwardID); err == nil && forward != nil {
|
||||
if quota, quotaErr := h.repo.AddTunnelQuotaUsage(forward.TunnelID, inFlow+outFlow, time.Now()); quotaErr == nil {
|
||||
h.enforceTunnelQuotaIfNeeded(forward.TunnelID, quota)
|
||||
}
|
||||
if quota, quotaErr := h.repo.AddUserQuotaUsage(userID, inFlow+outFlow, time.Now()); quotaErr == nil {
|
||||
h.enforceUserQuotaIfNeeded(userID, quota)
|
||||
}
|
||||
h.processPeerShareFlowFromForward(forwardID, nodeID, serviceName, item)
|
||||
|
||||
@@ -361,7 +359,7 @@ func (h *Handler) ensureUserTunnelForwardAllowed(userID int64, tunnelID int64, n
|
||||
if flowLimit < current {
|
||||
return errors.New("流量已超额,禁止开启转发")
|
||||
}
|
||||
if err := h.ensureTunnelForwardAllowedByQuota(tunnelID, now); err != nil {
|
||||
if err := h.ensureUserForwardAllowedByQuota(userID, now); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
|
||||
@@ -101,6 +101,7 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/user/update", h.userUpdate)
|
||||
mux.HandleFunc("/api/v1/user/delete", h.userDelete)
|
||||
mux.HandleFunc("/api/v1/user/reset", h.userResetFlow)
|
||||
mux.HandleFunc("/api/v1/user/quota/reset", h.userQuotaReset)
|
||||
mux.HandleFunc("/api/v1/user/groups", h.userGroups)
|
||||
mux.HandleFunc("/api/v1/config/get", h.getConfigByName)
|
||||
mux.HandleFunc("/api/v1/config/list", h.getConfigs)
|
||||
@@ -132,7 +133,6 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/tunnel/create", h.tunnelCreate)
|
||||
mux.HandleFunc("/api/v1/tunnel/get", h.tunnelGet)
|
||||
mux.HandleFunc("/api/v1/tunnel/update", h.tunnelUpdate)
|
||||
mux.HandleFunc("/api/v1/tunnel/quota/reset", h.tunnelQuotaReset)
|
||||
mux.HandleFunc("/api/v1/tunnel/delete", h.tunnelDelete)
|
||||
mux.HandleFunc("/api/v1/tunnel/diagnose", h.tunnelDiagnose)
|
||||
mux.HandleFunc("/api/v1/tunnel/diagnose/stream", h.tunnelDiagnoseStream)
|
||||
|
||||
@@ -136,7 +136,7 @@ func (h *Handler) runResetAndExpiryJob(now time.Time) {
|
||||
}
|
||||
|
||||
h.resetMonthlyFlow(now)
|
||||
h.resetTunnelQuotaWindows(now)
|
||||
h.resetUserQuotaWindows(now)
|
||||
h.disableExpiredUsers(now.UnixMilli())
|
||||
h.disableExpiredUserTunnels(now.UnixMilli())
|
||||
}
|
||||
|
||||
@@ -144,7 +144,7 @@ func TestRunResetAndExpiryJobResetsFlowAndDisablesExpiredRecords(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunResetAndExpiryJobResetsTunnelQuotaAndReEnablesTunnel(t *testing.T) {
|
||||
func TestRunResetAndExpiryJobResetsUserQuotaAndUnblocksUser(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "jobs-quota-reset.db")
|
||||
r, err := repo.Open(dbPath)
|
||||
if err != nil {
|
||||
@@ -157,29 +157,25 @@ func TestRunResetAndExpiryJobResetsTunnelQuotaAndReEnablesTunnel(t *testing.T) {
|
||||
nowMs := now.UnixMilli()
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(1, 'quota-reset-tunnel', 1.0, 1, 'tls', 1, ?, ?, 0, NULL, 0)
|
||||
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(2, 'quota-reset-user', 'x', 1, 0, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO tunnel_quota(tunnel_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time)
|
||||
VALUES(1, 10, 0, ?, ?, 20260311, 202603, 1, ?, '', ?, ?)
|
||||
INSERT INTO user_quota(user_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time)
|
||||
VALUES(2, 10, 0, ?, ?, 20260311, 202603, 1, ?, '', ?, ?)
|
||||
`, 11*int64(1024*1024*1024), 11*int64(1024*1024*1024), nowMs, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert tunnel quota: %v", err)
|
||||
t.Fatalf("insert user quota: %v", err)
|
||||
}
|
||||
|
||||
h.runResetAndExpiryJob(now)
|
||||
|
||||
tunnelStatus := mustQueryInt(t, r, `SELECT status FROM tunnel WHERE id = 1`)
|
||||
if tunnelStatus != 1 {
|
||||
t.Fatalf("expected tunnel re-enabled after quota reset, got %d", tunnelStatus)
|
||||
}
|
||||
dailyUsed := mustQueryInt(t, r, `SELECT daily_used_bytes FROM tunnel_quota WHERE tunnel_id = 1`)
|
||||
dailyUsed := mustQueryInt(t, r, `SELECT daily_used_bytes FROM user_quota WHERE user_id = 2`)
|
||||
if dailyUsed != 0 {
|
||||
t.Fatalf("expected daily quota usage reset, got %d", dailyUsed)
|
||||
}
|
||||
quotaDisabled := mustQueryInt(t, r, `SELECT disabled_by_quota FROM tunnel_quota WHERE tunnel_id = 1`)
|
||||
quotaDisabled := mustQueryInt(t, r, `SELECT disabled_by_quota FROM user_quota WHERE user_id = 2`)
|
||||
if quotaDisabled != 0 {
|
||||
t.Fatalf("expected quota disabled flag cleared, got %d", quotaDisabled)
|
||||
}
|
||||
|
||||
@@ -60,6 +60,12 @@ func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) {
|
||||
num := asInt(req["num"], 10)
|
||||
expTime := asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli())
|
||||
flowResetTime := asInt64(req["flowResetTime"], 1)
|
||||
dailyQuotaGB := asInt64(req["dailyQuotaGB"], 0)
|
||||
monthlyQuotaGB := asInt64(req["monthlyQuotaGB"], 0)
|
||||
if dailyQuotaGB < 0 || monthlyQuotaGB < 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("配额不能小于0"))
|
||||
return
|
||||
}
|
||||
roleID := 1
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
@@ -68,6 +74,22 @@ func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if dailyQuotaGB > 0 || monthlyQuotaGB > 0 {
|
||||
tx := h.repo.BeginTx()
|
||||
if tx == nil || tx.Error != nil {
|
||||
response.WriteJSON(w, response.Err(-2, "database unavailable"))
|
||||
return
|
||||
}
|
||||
defer func() { tx.Rollback() }()
|
||||
if err := h.repo.SaveUserQuotaConfigTx(tx, userID, dailyQuotaGB, monthlyQuotaGB, now); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := tx.Commit().Error; err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
groupIDs := asInt64Slice(req["groupIds"])
|
||||
if len(groupIDs) > 0 {
|
||||
@@ -131,6 +153,8 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
expTime := asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli())
|
||||
flowResetTime := asInt64(req["flowResetTime"], 1)
|
||||
status := asInt(req["status"], 1)
|
||||
_, hasDailyQuota := req["dailyQuotaGB"]
|
||||
_, hasMonthlyQuota := req["monthlyQuotaGB"]
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
pwd := asString(req["pwd"])
|
||||
@@ -147,6 +171,34 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
h.repo.PropagateUserFlowToTunnels(id, flow, num, expTime, flowResetTime)
|
||||
if hasDailyQuota || hasMonthlyQuota {
|
||||
dailyQuotaGB := asInt64(req["dailyQuotaGB"], 0)
|
||||
monthlyQuotaGB := asInt64(req["monthlyQuotaGB"], 0)
|
||||
if !(hasDailyQuota && hasMonthlyQuota) {
|
||||
if currentQuota, err := h.repo.GetUserQuotaView(id, time.Now()); err == nil && currentQuota != nil {
|
||||
if !hasDailyQuota {
|
||||
dailyQuotaGB = currentQuota.DailyLimitGB
|
||||
}
|
||||
if !hasMonthlyQuota {
|
||||
monthlyQuotaGB = currentQuota.MonthlyLimitGB
|
||||
}
|
||||
}
|
||||
}
|
||||
tx := h.repo.BeginTx()
|
||||
if tx == nil || tx.Error != nil {
|
||||
response.WriteJSON(w, response.Err(-2, "database unavailable"))
|
||||
return
|
||||
}
|
||||
defer func() { tx.Rollback() }()
|
||||
if err := h.repo.SaveUserQuotaConfigTx(tx, id, dailyQuotaGB, monthlyQuotaGB, now); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := tx.Commit().Error; err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if groupIDsRaw, ok := req["groupIds"]; ok {
|
||||
newGroupIDs := asInt64Slice(groupIDsRaw)
|
||||
@@ -485,8 +537,6 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
typeVal := asInt(req["type"], 1)
|
||||
flow := asInt64(req["flow"], 1)
|
||||
dailyQuotaGB := asInt64(req["dailyQuotaGB"], 0)
|
||||
monthlyQuotaGB := asInt64(req["monthlyQuotaGB"], 0)
|
||||
status := asInt(req["status"], 1)
|
||||
trafficRatio := asFloat(req["trafficRatio"], 1.0)
|
||||
inIP := asString(req["inIp"])
|
||||
@@ -582,10 +632,6 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
tunnelID := tunnel.ID
|
||||
if err := h.repo.SaveTunnelQuotaConfigTx(tx, tunnelID, dailyQuotaGB, monthlyQuotaGB, now); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
runtimeState.TunnelID = tunnelID
|
||||
var federationBindings []repo.FederationTunnelBinding
|
||||
var federationReleaseRefs []federationRuntimeReleaseRef
|
||||
@@ -697,8 +743,6 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
typeVal := asInt(req["type"], 1)
|
||||
dailyQuotaGB := asInt64(req["dailyQuotaGB"], 0)
|
||||
monthlyQuotaGB := asInt64(req["monthlyQuotaGB"], 0)
|
||||
ipPreference := asString(req["ipPreference"])
|
||||
localDomain := h.federationLocalDomain()
|
||||
|
||||
@@ -743,10 +787,6 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.repo.SaveTunnelQuotaConfigTx(tx, id, dailyQuotaGB, monthlyQuotaGB, now); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.repo.DeleteChainTunnelsByTunnelTx(tx, id); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
@@ -762,13 +802,25 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
newEntryNodeIDs := make([]int64, 0, len(runtimeState.InNodes))
|
||||
for _, in := range runtimeState.InNodes {
|
||||
if in.NodeID > 0 {
|
||||
newEntryNodeIDs = append(newEntryNodeIDs, in.NodeID)
|
||||
}
|
||||
}
|
||||
if err := h.validateTunnelEntryPortConflictsForNewEntries(id, oldEntryNodeIDs, newEntryNodeIDs); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
if err := tx.Commit().Error; err != nil {
|
||||
h.releaseFederationRuntimeRefs(federationReleaseRefs)
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
newEntryNodeIDs, _ := h.tunnelEntryNodeIDs(id)
|
||||
newEntryNodeIDs, _ = h.tunnelEntryNodeIDs(id)
|
||||
if !sameInt64Set(oldEntryNodeIDs, newEntryNodeIDs) {
|
||||
h.cleanupTunnelForwardRuntimesOnRemovedEntryNodes(id, oldEntryNodeIDs, newEntryNodeIDs)
|
||||
h.syncTunnelForwardsEntryPorts(id, newEntryNodeIDs)
|
||||
@@ -836,6 +888,55 @@ func pickForwardPortFromRecords(ports []forwardPortRecord) int {
|
||||
return min
|
||||
}
|
||||
|
||||
func forwardPortNodeIDs(ports []forwardPortRecord) []int64 {
|
||||
ids := make([]int64, 0, len(ports))
|
||||
for _, fp := range ports {
|
||||
if fp.NodeID <= 0 {
|
||||
continue
|
||||
}
|
||||
ids = append(ids, fp.NodeID)
|
||||
}
|
||||
return uniqueInt64s(ids)
|
||||
}
|
||||
|
||||
func (h *Handler) deleteForwardServicesOnNodeBatch(forward *forwardRecord, nodeID int64) error {
|
||||
if h == nil || forward == nil || nodeID <= 0 {
|
||||
return errors.New("invalid forward service cleanup context")
|
||||
}
|
||||
bases, err := h.forwardServiceBaseCandidates(forward)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(bases) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
names := make([]string, 0, len(bases)*3)
|
||||
seen := make(map[string]struct{}, len(bases)*3)
|
||||
appendName := func(name string) {
|
||||
if strings.TrimSpace(name) == "" {
|
||||
return
|
||||
}
|
||||
if _, ok := seen[name]; ok {
|
||||
return
|
||||
}
|
||||
seen[name] = struct{}{}
|
||||
names = append(names, name)
|
||||
}
|
||||
for _, base := range bases {
|
||||
appendName(base + "_tcp")
|
||||
appendName(base + "_udp")
|
||||
appendName(base)
|
||||
}
|
||||
if len(names) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
payload := map[string]interface{}{"services": names}
|
||||
_, err = h.sendNodeCommand(nodeID, "DeleteService", payload, false, true)
|
||||
return err
|
||||
}
|
||||
|
||||
func uniqueInt64s(input []int64) []int64 {
|
||||
if len(input) <= 1 {
|
||||
return input
|
||||
@@ -896,6 +997,52 @@ func (h *Handler) cleanupTunnelForwardRuntimesOnRemovedEntryNodes(tunnelID int64
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) validateTunnelEntryPortConflictsForNewEntries(tunnelID int64, oldEntryNodeIDs, newEntryNodeIDs []int64) error {
|
||||
if h == nil || h.repo == nil || tunnelID <= 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
addedNodeIDs := diffInt64s(newEntryNodeIDs, oldEntryNodeIDs)
|
||||
if len(addedNodeIDs) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
forwards, err := h.listForwardsByTunnel(tunnelID)
|
||||
if err != nil || len(forwards) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
for i := range forwards {
|
||||
f := &forwards[i]
|
||||
if f == nil {
|
||||
continue
|
||||
}
|
||||
oldPorts, portsErr := h.listForwardPorts(f.ID)
|
||||
if portsErr != nil {
|
||||
continue
|
||||
}
|
||||
port := pickForwardPortFromRecords(oldPorts)
|
||||
if port <= 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
for _, nodeID := range addedNodeIDs {
|
||||
node, nodeErr := h.getNodeRecord(nodeID)
|
||||
if nodeErr != nil {
|
||||
continue
|
||||
}
|
||||
if err := validateLocalNodePort(node, port); err != nil {
|
||||
return fmt.Errorf("转发 %s 入口端口冲突: %w", f.Name, err)
|
||||
}
|
||||
if err := h.validateForwardPortAvailability(node, port, f.ID); err != nil {
|
||||
return fmt.Errorf("转发 %s 入口端口冲突: %w", f.Name, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) syncTunnelForwardsEntryPorts(tunnelID int64, entryNodeIDs []int64) {
|
||||
if h == nil || h.repo == nil || tunnelID <= 0 {
|
||||
return
|
||||
@@ -1007,16 +1154,24 @@ func (h *Handler) tunnelBatchDelete(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
success := 0
|
||||
fail := 0
|
||||
failures := make([]batchFailureDetail, 0)
|
||||
for _, id := range ids {
|
||||
tunnelName, _ := h.repo.GetTunnelName(id)
|
||||
if _, err := h.getTunnelRecord(id); err != nil {
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, id, tunnelName, err)
|
||||
continue
|
||||
}
|
||||
h.cleanupTunnelRuntime(id)
|
||||
h.cleanupFederationRuntime(id)
|
||||
if err := h.deleteTunnelByID(id); err != nil {
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, id, tunnelName, err)
|
||||
} else {
|
||||
success++
|
||||
}
|
||||
}
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": success, "failCount": fail}))
|
||||
response.WriteJSON(w, response.OK(batchOperationResult{SuccessCount: success, FailCount: fail, Failures: failures}))
|
||||
}
|
||||
|
||||
func (h *Handler) reconstructTunnelState(tunnelID int64) (*tunnelCreateState, error) {
|
||||
@@ -1152,6 +1307,60 @@ func (h *Handler) redeployTunnelAndForwards(tunnelID int64) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
type batchFailureDetail struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name,omitempty"`
|
||||
Reason string `json:"reason"`
|
||||
}
|
||||
|
||||
type batchOperationResult struct {
|
||||
SuccessCount int `json:"successCount"`
|
||||
FailCount int `json:"failCount"`
|
||||
Failures []batchFailureDetail `json:"failures,omitempty"`
|
||||
}
|
||||
|
||||
func appendBatchFailure(failures []batchFailureDetail, id int64, name string, err error) []batchFailureDetail {
|
||||
reason := normalizeBatchFailureReason(errString(err))
|
||||
if reason == "" {
|
||||
reason = "未知错误"
|
||||
}
|
||||
return append(failures, batchFailureDetail{
|
||||
ID: id,
|
||||
Name: strings.TrimSpace(name),
|
||||
Reason: reason,
|
||||
})
|
||||
}
|
||||
|
||||
func appendBatchFailureReason(failures []batchFailureDetail, id int64, name, reason string) []batchFailureDetail {
|
||||
normalized := normalizeBatchFailureReason(reason)
|
||||
if normalized == "" {
|
||||
normalized = "未知错误"
|
||||
}
|
||||
return append(failures, batchFailureDetail{
|
||||
ID: id,
|
||||
Name: strings.TrimSpace(name),
|
||||
Reason: normalized,
|
||||
})
|
||||
}
|
||||
|
||||
func errString(err error) string {
|
||||
if err == nil {
|
||||
return ""
|
||||
}
|
||||
return err.Error()
|
||||
}
|
||||
|
||||
func normalizeBatchFailureReason(reason string) string {
|
||||
trimmed := strings.TrimSpace(reason)
|
||||
if trimmed == "" {
|
||||
return ""
|
||||
}
|
||||
if strings.EqualFold(trimmed, errForwardNotFound.Error()) {
|
||||
return "转发不存在"
|
||||
}
|
||||
return trimmed
|
||||
}
|
||||
|
||||
func (h *Handler) tunnelBatchRedeploy(w http.ResponseWriter, r *http.Request) {
|
||||
ids := idsFromBody(r, w)
|
||||
if ids == nil {
|
||||
@@ -1159,14 +1368,26 @@ func (h *Handler) tunnelBatchRedeploy(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
success := 0
|
||||
fail := 0
|
||||
failures := make([]batchFailureDetail, 0)
|
||||
for _, tunnelID := range ids {
|
||||
tunnel, tunnelErr := h.getTunnelRecord(tunnelID)
|
||||
tunnelName := ""
|
||||
if tunnelErr == nil && tunnel != nil {
|
||||
tunnelName, _ = h.repo.GetTunnelName(tunnelID)
|
||||
}
|
||||
if tunnelErr != nil {
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, tunnelID, tunnelName, tunnelErr)
|
||||
continue
|
||||
}
|
||||
if err := h.redeployTunnelAndForwards(tunnelID); err != nil {
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, tunnelID, tunnelName, err)
|
||||
continue
|
||||
}
|
||||
success++
|
||||
}
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": success, "failCount": fail}))
|
||||
response.WriteJSON(w, response.OK(batchOperationResult{SuccessCount: success, FailCount: fail, Failures: failures}))
|
||||
}
|
||||
|
||||
func (h *Handler) userTunnelAssign(w http.ResponseWriter, r *http.Request) {
|
||||
@@ -1311,10 +1532,6 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
if tunnel.Status != 1 {
|
||||
if reason, quotaErr := h.tunnelQuotaBlockReason(tunnelID, time.Now().UnixMilli()); quotaErr == nil && reason != "" {
|
||||
response.WriteJSON(w, response.ErrDefault(reason))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.ErrDefault("隧道已禁用,无法创建转发"))
|
||||
return
|
||||
}
|
||||
@@ -1436,10 +1653,6 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
if tunnel.Status != 1 {
|
||||
if reason, quotaErr := h.tunnelQuotaBlockReason(tunnelID, time.Now().UnixMilli()); quotaErr == nil && reason != "" {
|
||||
response.WriteJSON(w, response.ErrDefault(reason))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.ErrDefault("隧道已禁用,无法更新转发"))
|
||||
return
|
||||
}
|
||||
@@ -1496,6 +1709,17 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("多入口隧道的转发不支持自定义监听IP"))
|
||||
return
|
||||
}
|
||||
// When switching tunnels, entry nodes / service base may change. We must clean up old
|
||||
// listeners on nodes that will be removed, otherwise the old ports keep listening.
|
||||
tunnelChanged := tunnelID != forward.TunnelID
|
||||
oldNodeIDs := forwardPortNodeIDs(oldPorts)
|
||||
newNodeIDs := uniqueInt64s(fwdEntryNodes)
|
||||
var removedNodeIDs []int64
|
||||
var keptNodeIDs []int64
|
||||
if tunnelChanged {
|
||||
removedNodeIDs = diffInt64s(oldNodeIDs, newNodeIDs)
|
||||
keptNodeIDs = diffInt64s(oldNodeIDs, removedNodeIDs)
|
||||
}
|
||||
for _, nodeID := range fwdEntryNodes {
|
||||
node, nodeErr := h.getNodeRecord(nodeID)
|
||||
if nodeErr != nil {
|
||||
@@ -1537,12 +1761,41 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
warnings, err := h.syncForwardServicesWithWarnings(updatedForward, "UpdateService", true)
|
||||
warnings := make([]string, 0)
|
||||
if tunnelChanged && len(keptNodeIDs) > 0 {
|
||||
for _, nodeID := range keptNodeIDs {
|
||||
if delErr := h.deleteForwardServicesOnNodeBatch(forward, nodeID); delErr != nil {
|
||||
nodeLabel := fmt.Sprintf("%d", nodeID)
|
||||
if n, nErr := h.getNodeRecord(nodeID); nErr == nil && n != nil && strings.TrimSpace(n.Name) != "" {
|
||||
nodeLabel = strings.TrimSpace(n.Name)
|
||||
}
|
||||
warnings = append(warnings, fmt.Sprintf("节点 %s 清理旧转发监听失败: %v", nodeLabel, delErr))
|
||||
}
|
||||
}
|
||||
time.Sleep(tunnelServiceBindRetryDelay)
|
||||
}
|
||||
|
||||
syncWarnings, err := h.syncForwardServicesWithWarnings(updatedForward, "UpdateService", true)
|
||||
if err != nil {
|
||||
h.rollbackForwardMutation(forward, oldPorts)
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
warnings = append(warnings, syncWarnings...)
|
||||
|
||||
// Best-effort cleanup for old entry nodes after a successful tunnel switch.
|
||||
// Avoid cleaning nodes that are still used by the updated forward.
|
||||
if tunnelChanged && len(removedNodeIDs) > 0 {
|
||||
for _, nodeID := range removedNodeIDs {
|
||||
if delErr := h.deleteForwardServicesOnNodeBatch(forward, nodeID); delErr != nil {
|
||||
nodeLabel := fmt.Sprintf("%d", nodeID)
|
||||
if n, nErr := h.getNodeRecord(nodeID); nErr == nil && n != nil && strings.TrimSpace(n.Name) != "" {
|
||||
nodeLabel = strings.TrimSpace(n.Name)
|
||||
}
|
||||
warnings = append(warnings, fmt.Sprintf("节点 %s 清理旧隧道残留服务失败: %v", nodeLabel, delErr))
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(warnings) > 0 {
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{"warnings": warnings}))
|
||||
return
|
||||
@@ -1685,23 +1938,27 @@ func (h *Handler) forwardBatchDelete(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
s := 0
|
||||
f := 0
|
||||
failures := make([]batchFailureDetail, 0)
|
||||
for _, id := range ids {
|
||||
forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id)
|
||||
if accessErr != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, "", accessErr)
|
||||
continue
|
||||
}
|
||||
if err := h.controlForwardServices(forward, "DeleteService", true); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
continue
|
||||
}
|
||||
if err := h.deleteForwardByID(id); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
} else {
|
||||
s++
|
||||
}
|
||||
}
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": s, "failCount": f}))
|
||||
response.WriteJSON(w, response.OK(batchOperationResult{SuccessCount: s, FailCount: f, Failures: failures}))
|
||||
}
|
||||
|
||||
func (h *Handler) forwardBatchPause(w http.ResponseWriter, r *http.Request) {
|
||||
@@ -1716,23 +1973,27 @@ func (h *Handler) forwardBatchPause(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
s := 0
|
||||
f := 0
|
||||
failures := make([]batchFailureDetail, 0)
|
||||
for _, id := range ids {
|
||||
forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id)
|
||||
if accessErr != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, "", accessErr)
|
||||
continue
|
||||
}
|
||||
if err := h.controlForwardServices(forward, "PauseService", false); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
continue
|
||||
}
|
||||
if err := h.repo.UpdateForwardStatus(id, 0, time.Now().UnixMilli()); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
} else {
|
||||
s++
|
||||
}
|
||||
}
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": s, "failCount": f}))
|
||||
response.WriteJSON(w, response.OK(batchOperationResult{SuccessCount: s, FailCount: f, Failures: failures}))
|
||||
}
|
||||
|
||||
func (h *Handler) forwardBatchResume(w http.ResponseWriter, r *http.Request) {
|
||||
@@ -1748,27 +2009,32 @@ func (h *Handler) forwardBatchResume(w http.ResponseWriter, r *http.Request) {
|
||||
s := 0
|
||||
f := 0
|
||||
now := time.Now().UnixMilli()
|
||||
failures := make([]batchFailureDetail, 0)
|
||||
for _, id := range ids {
|
||||
forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id)
|
||||
if accessErr != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, "", accessErr)
|
||||
continue
|
||||
}
|
||||
if err := h.ensureUserTunnelForwardAllowed(forward.UserID, forward.TunnelID, now); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
continue
|
||||
}
|
||||
if err := h.controlForwardServices(forward, "ResumeService", false); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
continue
|
||||
}
|
||||
if err := h.repo.UpdateForwardStatus(id, 1, now); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
} else {
|
||||
s++
|
||||
}
|
||||
}
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": s, "failCount": f}))
|
||||
response.WriteJSON(w, response.OK(batchOperationResult{SuccessCount: s, FailCount: f, Failures: failures}))
|
||||
}
|
||||
|
||||
func (h *Handler) forwardBatchRedeploy(w http.ResponseWriter, r *http.Request) {
|
||||
@@ -1783,19 +2049,22 @@ func (h *Handler) forwardBatchRedeploy(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
s := 0
|
||||
f := 0
|
||||
failures := make([]batchFailureDetail, 0)
|
||||
for _, id := range ids {
|
||||
forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id)
|
||||
if accessErr != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, "", accessErr)
|
||||
continue
|
||||
}
|
||||
if err := h.syncForwardServices(forward, "UpdateService", true); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
} else {
|
||||
s++
|
||||
}
|
||||
}
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": s, "failCount": f}))
|
||||
response.WriteJSON(w, response.OK(batchOperationResult{SuccessCount: s, FailCount: f, Failures: failures}))
|
||||
}
|
||||
|
||||
func (h *Handler) forwardBatchChangeTunnel(w http.ResponseWriter, r *http.Request) {
|
||||
@@ -1827,6 +2096,7 @@ func (h *Handler) forwardBatchChangeTunnel(w http.ResponseWriter, r *http.Reques
|
||||
}
|
||||
success := 0
|
||||
fail := 0
|
||||
failures := make([]batchFailureDetail, 0)
|
||||
for _, id := range req.ForwardIDs {
|
||||
if id <= 0 {
|
||||
continue
|
||||
@@ -1834,20 +2104,30 @@ func (h *Handler) forwardBatchChangeTunnel(w http.ResponseWriter, r *http.Reques
|
||||
forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id)
|
||||
if accessErr != nil {
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, id, "", accessErr)
|
||||
continue
|
||||
}
|
||||
if forward.TunnelID == req.TargetTunnelID {
|
||||
fail++
|
||||
failures = appendBatchFailureReason(failures, id, forward.Name, "规则已在目标隧道中")
|
||||
continue
|
||||
}
|
||||
oldPorts, listPortsErr := h.listForwardPorts(id)
|
||||
if listPortsErr != nil {
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, listPortsErr)
|
||||
continue
|
||||
}
|
||||
if len(oldPorts) == 0 {
|
||||
fail++
|
||||
failures = appendBatchFailureReason(failures, id, forward.Name, "转发入口端口不存在")
|
||||
continue
|
||||
}
|
||||
oldNodeIDs := forwardPortNodeIDs(oldPorts)
|
||||
port := h.repo.GetMinForwardPort(id)
|
||||
if err := h.repo.UpdateForwardTunnel(id, req.TargetTunnelID, time.Now().UnixMilli()); err != nil {
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
continue
|
||||
}
|
||||
p := 0
|
||||
@@ -1858,40 +2138,62 @@ func (h *Handler) forwardBatchChangeTunnel(w http.ResponseWriter, r *http.Reques
|
||||
p = h.pickTunnelPort(req.TargetTunnelID)
|
||||
}
|
||||
bctEntryNodes, _ := h.tunnelEntryNodeIDs(req.TargetTunnelID)
|
||||
newNodeIDs := uniqueInt64s(bctEntryNodes)
|
||||
removedNodeIDs := diffInt64s(oldNodeIDs, newNodeIDs)
|
||||
keptNodeIDs := diffInt64s(oldNodeIDs, removedNodeIDs)
|
||||
portRangeOk := true
|
||||
var portRangeErr error
|
||||
for _, nid := range bctEntryNodes {
|
||||
nd, ndErr := h.getNodeRecord(nid)
|
||||
if ndErr != nil {
|
||||
portRangeErr = ndErr
|
||||
continue
|
||||
}
|
||||
if validateRemoteNodePort(nd, p) != nil {
|
||||
if validateErr := validateRemoteNodePort(nd, p); validateErr != nil {
|
||||
portRangeOk = false
|
||||
portRangeErr = validateErr
|
||||
break
|
||||
}
|
||||
}
|
||||
if !portRangeOk {
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, portRangeErr)
|
||||
h.rollbackForwardMutation(forward, oldPorts)
|
||||
continue
|
||||
}
|
||||
if err := h.replaceForwardPorts(id, req.TargetTunnelID, p, ""); err != nil {
|
||||
h.rollbackForwardMutation(forward, oldPorts)
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
continue
|
||||
}
|
||||
updatedForward, fetchErr := h.getForwardRecord(id)
|
||||
if fetchErr != nil {
|
||||
h.rollbackForwardMutation(forward, oldPorts)
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, fetchErr)
|
||||
continue
|
||||
}
|
||||
if len(keptNodeIDs) > 0 {
|
||||
for _, nodeID := range keptNodeIDs {
|
||||
_ = h.deleteForwardServicesOnNodeBatch(forward, nodeID)
|
||||
}
|
||||
time.Sleep(tunnelServiceBindRetryDelay)
|
||||
}
|
||||
if err := h.syncForwardServices(updatedForward, "UpdateService", true); err != nil {
|
||||
h.rollbackForwardMutation(forward, oldPorts)
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
continue
|
||||
}
|
||||
if len(removedNodeIDs) > 0 {
|
||||
for _, nodeID := range removedNodeIDs {
|
||||
_ = h.deleteForwardServicesOnNodeBatch(forward, nodeID)
|
||||
}
|
||||
}
|
||||
success++
|
||||
}
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": success, "failCount": fail}))
|
||||
response.WriteJSON(w, response.OK(batchOperationResult{SuccessCount: success, FailCount: fail, Failures: failures}))
|
||||
}
|
||||
|
||||
func (h *Handler) speedLimitCreate(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
@@ -1,140 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func isTunnelQuotaExceeded(view *model.TunnelQuotaView) bool {
|
||||
if view == nil {
|
||||
return false
|
||||
}
|
||||
if view.DailyLimitGB > 0 && view.DailyUsedBytes >= view.DailyLimitGB*bytesPerGB {
|
||||
return true
|
||||
}
|
||||
if view.MonthlyLimitGB > 0 && view.MonthlyUsedBytes >= view.MonthlyLimitGB*bytesPerGB {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (h *Handler) tunnelQuotaBlockReason(tunnelID int64, now int64) (string, error) {
|
||||
if h == nil || h.repo == nil || tunnelID <= 0 {
|
||||
return "", nil
|
||||
}
|
||||
quota, err := h.repo.GetTunnelQuotaView(tunnelID, time.UnixMilli(now))
|
||||
if err != nil || quota == nil {
|
||||
return "", err
|
||||
}
|
||||
if quota.DisabledByQuota == 1 || isTunnelQuotaExceeded(quota) {
|
||||
return "该隧道流量配额已超额,禁止开启转发", nil
|
||||
}
|
||||
return "", nil
|
||||
}
|
||||
|
||||
func (h *Handler) enforceTunnelQuotaIfNeeded(tunnelID int64, quota *model.TunnelQuotaView) {
|
||||
if h == nil || h.repo == nil || tunnelID <= 0 || quota == nil {
|
||||
return
|
||||
}
|
||||
if quota.DisabledByQuota == 1 || !isTunnelQuotaExceeded(quota) {
|
||||
return
|
||||
}
|
||||
|
||||
forwards, err := h.listForwardsByTunnel(tunnelID)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
pausedIDs := make([]int64, 0, len(forwards))
|
||||
now := time.Now().UnixMilli()
|
||||
for i := range forwards {
|
||||
if forwards[i].Status != 1 {
|
||||
continue
|
||||
}
|
||||
if err := h.controlForwardServices(&forwards[i], "PauseService", false); err != nil {
|
||||
continue
|
||||
}
|
||||
if err := h.repo.UpdateForwardStatus(forwards[i].ID, 0, now); err != nil {
|
||||
continue
|
||||
}
|
||||
pausedIDs = append(pausedIDs, forwards[i].ID)
|
||||
}
|
||||
_ = h.repo.UpdateTunnelStatus(tunnelID, 0, now)
|
||||
_ = h.repo.MarkTunnelQuotaDisabled(tunnelID, pausedIDs, now)
|
||||
}
|
||||
|
||||
func (h *Handler) applyTunnelQuotaRelease(release *repo.TunnelQuotaRelease, now int64) {
|
||||
if h == nil || h.repo == nil || release == nil || release.TunnelID <= 0 || !release.EnableTunnel {
|
||||
return
|
||||
}
|
||||
_ = h.repo.UpdateTunnelStatus(release.TunnelID, 1, now)
|
||||
for _, forwardID := range release.ForwardIDs {
|
||||
forward, err := h.getForwardRecord(forwardID)
|
||||
if err != nil || forward == nil {
|
||||
continue
|
||||
}
|
||||
if err := h.ensureUserTunnelForwardAllowed(forward.UserID, forward.TunnelID, now); err != nil {
|
||||
continue
|
||||
}
|
||||
if err := h.controlForwardServices(forward, "ResumeService", false); err != nil {
|
||||
continue
|
||||
}
|
||||
_ = h.repo.UpdateForwardStatus(forwardID, 1, now)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) resetTunnelQuotaWindows(now time.Time) {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
releases, err := h.repo.RollTunnelQuotaWindows(now)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
nowMs := now.UnixMilli()
|
||||
for i := range releases {
|
||||
h.applyTunnelQuotaRelease(&releases[i], nowMs)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) tunnelQuotaReset(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
var req struct {
|
||||
TunnelID int64 `json:"tunnelId"`
|
||||
Scope string `json:"scope"`
|
||||
}
|
||||
if err := decodeJSON(r.Body, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
if req.TunnelID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("隧道ID不能为空"))
|
||||
return
|
||||
}
|
||||
release, err := h.repo.ResetTunnelQuotaUsage(req.TunnelID, req.Scope, time.Now())
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
nowMs := time.Now().UnixMilli()
|
||||
h.applyTunnelQuotaRelease(release, nowMs)
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func (h *Handler) ensureTunnelForwardAllowedByQuota(tunnelID int64, now int64) error {
|
||||
reason, err := h.tunnelQuotaBlockReason(tunnelID, now)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if reason != "" {
|
||||
return errors.New(reason)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,139 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func isUserQuotaExceeded(view *model.UserQuotaView) bool {
|
||||
if view == nil {
|
||||
return false
|
||||
}
|
||||
if view.DailyLimitGB > 0 && view.DailyUsedBytes >= view.DailyLimitGB*bytesPerGB {
|
||||
return true
|
||||
}
|
||||
if view.MonthlyLimitGB > 0 && view.MonthlyUsedBytes >= view.MonthlyLimitGB*bytesPerGB {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (h *Handler) userQuotaBlockReason(userID int64, now int64) (string, error) {
|
||||
if h == nil || h.repo == nil || userID <= 0 {
|
||||
return "", nil
|
||||
}
|
||||
quota, err := h.repo.GetUserQuotaView(userID, time.UnixMilli(now))
|
||||
if err != nil || quota == nil {
|
||||
return "", err
|
||||
}
|
||||
if quota.DisabledByQuota == 1 || isUserQuotaExceeded(quota) {
|
||||
return "该用户流量配额已超额,禁止开启转发", nil
|
||||
}
|
||||
return "", nil
|
||||
}
|
||||
|
||||
func (h *Handler) enforceUserQuotaIfNeeded(userID int64, quota *model.UserQuotaView) {
|
||||
if h == nil || h.repo == nil || userID <= 0 || quota == nil {
|
||||
return
|
||||
}
|
||||
if quota.DisabledByQuota == 1 || !isUserQuotaExceeded(quota) {
|
||||
return
|
||||
}
|
||||
|
||||
forwards, err := h.listActiveForwardsByUser(userID)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
pausedIDs := make([]int64, 0, len(forwards))
|
||||
now := time.Now().UnixMilli()
|
||||
for i := range forwards {
|
||||
forward := &forwards[i]
|
||||
if forward.Status != 1 {
|
||||
continue
|
||||
}
|
||||
if err := h.controlForwardServices(forward, "PauseService", false); err != nil {
|
||||
continue
|
||||
}
|
||||
if err := h.repo.UpdateForwardStatus(forward.ID, 0, now); err != nil {
|
||||
continue
|
||||
}
|
||||
pausedIDs = append(pausedIDs, forward.ID)
|
||||
}
|
||||
_ = h.repo.MarkUserQuotaDisabled(userID, pausedIDs, now)
|
||||
}
|
||||
|
||||
func (h *Handler) applyUserQuotaRelease(release *repo.UserQuotaRelease, now int64) {
|
||||
if h == nil || h.repo == nil || release == nil || release.UserID <= 0 || !release.UnblockUser {
|
||||
return
|
||||
}
|
||||
for _, forwardID := range release.ForwardIDs {
|
||||
forward, err := h.getForwardRecord(forwardID)
|
||||
if err != nil || forward == nil {
|
||||
continue
|
||||
}
|
||||
if err := h.ensureUserTunnelForwardAllowed(forward.UserID, forward.TunnelID, now); err != nil {
|
||||
continue
|
||||
}
|
||||
if err := h.controlForwardServices(forward, "ResumeService", false); err != nil {
|
||||
continue
|
||||
}
|
||||
_ = h.repo.UpdateForwardStatus(forwardID, 1, now)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) resetUserQuotaWindows(now time.Time) {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
releases, err := h.repo.RollUserQuotaWindows(now)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
nowMs := now.UnixMilli()
|
||||
for i := range releases {
|
||||
h.applyUserQuotaRelease(&releases[i], nowMs)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) userQuotaReset(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
var req struct {
|
||||
UserID int64 `json:"userId"`
|
||||
Scope string `json:"scope"`
|
||||
}
|
||||
if err := decodeJSON(r.Body, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
if req.UserID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("用户ID不能为空"))
|
||||
return
|
||||
}
|
||||
release, err := h.repo.ResetUserQuotaUsage(req.UserID, req.Scope, time.Now())
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
nowMs := time.Now().UnixMilli()
|
||||
h.applyUserQuotaRelease(release, nowMs)
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func (h *Handler) ensureUserForwardAllowedByQuota(userID int64, now int64) error {
|
||||
reason, err := h.userQuotaBlockReason(userID, now)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if reason != "" {
|
||||
return errors.New(reason)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -129,8 +129,8 @@ type Tunnel struct {
|
||||
|
||||
func (Tunnel) TableName() string { return "tunnel" }
|
||||
|
||||
type TunnelQuota struct {
|
||||
TunnelID int64 `gorm:"column:tunnel_id;primaryKey"`
|
||||
type UserQuota struct {
|
||||
UserID int64 `gorm:"column:user_id;primaryKey"`
|
||||
DailyLimitGB int64 `gorm:"column:daily_limit_gb;not null;default:0"`
|
||||
MonthlyLimitGB int64 `gorm:"column:monthly_limit_gb;not null;default:0"`
|
||||
DailyUsedBytes int64 `gorm:"column:daily_used_bytes;not null;default:0"`
|
||||
@@ -144,7 +144,7 @@ type TunnelQuota struct {
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
}
|
||||
|
||||
func (TunnelQuota) TableName() string { return "tunnel_quota" }
|
||||
func (UserQuota) TableName() string { return "user_quota" }
|
||||
|
||||
type ChainTunnel struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
@@ -338,19 +338,23 @@ type BackupData struct {
|
||||
}
|
||||
|
||||
type UserBackup struct {
|
||||
ID int64 `json:"id"`
|
||||
User string `json:"user"`
|
||||
Pwd string `json:"pwd"`
|
||||
RoleID int `json:"roleId"`
|
||||
ExpTime int64 `json:"expTime"`
|
||||
Flow int64 `json:"flow"`
|
||||
InFlow int64 `json:"inFlow"`
|
||||
OutFlow int64 `json:"outFlow"`
|
||||
FlowResetTime int64 `json:"flowResetTime"`
|
||||
Num int `json:"num"`
|
||||
CreatedTime int64 `json:"createdTime"`
|
||||
UpdatedTime int64 `json:"updatedTime,omitempty"`
|
||||
Status int `json:"status"`
|
||||
ID int64 `json:"id"`
|
||||
User string `json:"user"`
|
||||
Pwd string `json:"pwd"`
|
||||
RoleID int `json:"roleId"`
|
||||
ExpTime int64 `json:"expTime"`
|
||||
Flow int64 `json:"flow"`
|
||||
InFlow int64 `json:"inFlow"`
|
||||
OutFlow int64 `json:"outFlow"`
|
||||
FlowResetTime int64 `json:"flowResetTime"`
|
||||
DailyQuotaGB int64 `json:"dailyQuotaGB,omitempty"`
|
||||
MonthlyQuotaGB int64 `json:"monthlyQuotaGB,omitempty"`
|
||||
DisabledByQuota int `json:"disabledByQuota,omitempty"`
|
||||
QuotaDisabledAt int64 `json:"quotaDisabledAt,omitempty"`
|
||||
Num int `json:"num"`
|
||||
CreatedTime int64 `json:"createdTime"`
|
||||
UpdatedTime int64 `json:"updatedTime,omitempty"`
|
||||
Status int `json:"status"`
|
||||
}
|
||||
|
||||
type NodeBackup struct {
|
||||
@@ -383,23 +387,19 @@ type NodeBackup struct {
|
||||
}
|
||||
|
||||
type TunnelBackup struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
TrafficRatio float64 `json:"trafficRatio"`
|
||||
Type int `json:"type"`
|
||||
Protocol string `json:"protocol"`
|
||||
Flow int64 `json:"flow"`
|
||||
CreatedTime int64 `json:"createdTime"`
|
||||
UpdatedTime int64 `json:"updatedTime"`
|
||||
Status int `json:"status"`
|
||||
InIP string `json:"inIp,omitempty"`
|
||||
Inx int `json:"inx"`
|
||||
IPPreference string `json:"ipPreference,omitempty"`
|
||||
DailyQuotaGB int64 `json:"dailyQuotaGB,omitempty"`
|
||||
MonthlyQuotaGB int64 `json:"monthlyQuotaGB,omitempty"`
|
||||
DisabledByQuota int `json:"disabledByQuota,omitempty"`
|
||||
QuotaDisabledAt int64 `json:"quotaDisabledAt,omitempty"`
|
||||
ChainTunnels []ChainTunnelBackup `json:"chainTunnels,omitempty"`
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
TrafficRatio float64 `json:"trafficRatio"`
|
||||
Type int `json:"type"`
|
||||
Protocol string `json:"protocol"`
|
||||
Flow int64 `json:"flow"`
|
||||
CreatedTime int64 `json:"createdTime"`
|
||||
UpdatedTime int64 `json:"updatedTime"`
|
||||
Status int `json:"status"`
|
||||
InIP string `json:"inIp,omitempty"`
|
||||
Inx int `json:"inx"`
|
||||
IPPreference string `json:"ipPreference,omitempty"`
|
||||
ChainTunnels []ChainTunnelBackup `json:"chainTunnels,omitempty"`
|
||||
}
|
||||
|
||||
type ChainTunnelBackup struct {
|
||||
@@ -537,8 +537,8 @@ type TunnelRecord struct {
|
||||
TrafficRatio float64
|
||||
}
|
||||
|
||||
type TunnelQuotaView struct {
|
||||
TunnelID int64
|
||||
type UserQuotaView struct {
|
||||
UserID int64
|
||||
DailyLimitGB int64
|
||||
MonthlyLimitGB int64
|
||||
DailyUsedBytes int64
|
||||
|
||||
@@ -161,13 +161,13 @@ func (r *Repository) Close() error {
|
||||
func autoMigrateAll(db *gorm.DB) error {
|
||||
models := []interface{}{
|
||||
&model.User{},
|
||||
&model.UserQuota{},
|
||||
&model.Forward{},
|
||||
&model.ForwardPort{},
|
||||
&model.Node{},
|
||||
&model.SpeedLimit{},
|
||||
&model.StatisticsFlow{},
|
||||
&model.Tunnel{},
|
||||
&model.TunnelQuota{},
|
||||
&model.ChainTunnel{},
|
||||
&model.UserTunnel{},
|
||||
&model.TunnelGroup{},
|
||||
@@ -664,16 +664,33 @@ func (r *Repository) ListUsers() ([]map[string]interface{}, error) {
|
||||
if err := r.db.Where("role_id != ?", 0).Order("id DESC").Find(&users).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
userIDs := make([]int64, 0, len(users))
|
||||
for _, u := range users {
|
||||
userIDs = append(userIDs, u.ID)
|
||||
}
|
||||
quotaMap, err := r.ListUserQuotaViewsByUserIDs(userIDs, time.Now())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items := make([]map[string]interface{}, 0, len(users))
|
||||
for _, u := range users {
|
||||
items = append(items, map[string]interface{}{
|
||||
item := map[string]interface{}{
|
||||
"id": u.ID, "user": u.User, "name": u.User,
|
||||
"roleId": u.RoleID, "status": u.Status,
|
||||
"flow": u.Flow, "num": u.Num, "expTime": u.ExpTime,
|
||||
"flowResetTime": u.FlowResetTime, "createdTime": u.CreatedTime,
|
||||
"updatedTime": nullableInt64(u.UpdatedTime),
|
||||
"inFlow": u.InFlow, "outFlow": u.OutFlow,
|
||||
})
|
||||
}
|
||||
if quota := quotaMap[u.ID]; quota != nil {
|
||||
item["dailyQuotaGB"] = quota.DailyLimitGB
|
||||
item["monthlyQuotaGB"] = quota.MonthlyLimitGB
|
||||
item["dailyUsedBytes"] = quota.DailyUsedBytes
|
||||
item["monthlyUsedBytes"] = quota.MonthlyUsedBytes
|
||||
item["disabledByQuota"] = quota.DisabledByQuota
|
||||
item["quotaDisabledAt"] = quota.DisabledAt
|
||||
}
|
||||
items = append(items, item)
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
@@ -958,7 +975,6 @@ func (r *Repository) ListTunnels() ([]map[string]interface{}, error) {
|
||||
|
||||
tunnelMap := make(map[int64]map[string]interface{})
|
||||
orderedIDs := make([]int64, 0, len(tunnels))
|
||||
tunnelIDs := make([]int64, 0, len(tunnels))
|
||||
|
||||
for _, t := range tunnels {
|
||||
tunnelMap[t.ID] = map[string]interface{}{
|
||||
@@ -972,24 +988,6 @@ func (r *Repository) ListTunnels() ([]map[string]interface{}, error) {
|
||||
"chainNodes": make([][]map[string]interface{}, 0),
|
||||
}
|
||||
orderedIDs = append(orderedIDs, t.ID)
|
||||
tunnelIDs = append(tunnelIDs, t.ID)
|
||||
}
|
||||
|
||||
quotaMap, err := r.ListTunnelQuotaViewsByTunnelIDs(tunnelIDs, time.Now())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for tunnelID, quota := range quotaMap {
|
||||
item := tunnelMap[tunnelID]
|
||||
if item == nil || quota == nil {
|
||||
continue
|
||||
}
|
||||
item["dailyQuotaGB"] = quota.DailyLimitGB
|
||||
item["monthlyQuotaGB"] = quota.MonthlyLimitGB
|
||||
item["dailyUsedBytes"] = quota.DailyUsedBytes
|
||||
item["monthlyUsedBytes"] = quota.MonthlyUsedBytes
|
||||
item["disabledByQuota"] = quota.DisabledByQuota
|
||||
item["quotaDisabledAt"] = quota.DisabledAt
|
||||
}
|
||||
|
||||
// Build node IP map
|
||||
@@ -1814,6 +1812,14 @@ func (r *Repository) exportUsers() ([]model.UserBackup, error) {
|
||||
if err := r.db.Order("id ASC").Find(&users).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
userIDs := make([]int64, 0, len(users))
|
||||
for _, u := range users {
|
||||
userIDs = append(userIDs, u.ID)
|
||||
}
|
||||
quotaMap, err := r.ListUserQuotaViewsByUserIDs(userIDs, time.Now())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out := make([]model.UserBackup, 0, len(users))
|
||||
for _, u := range users {
|
||||
b := model.UserBackup{
|
||||
@@ -1822,6 +1828,12 @@ func (r *Repository) exportUsers() ([]model.UserBackup, error) {
|
||||
FlowResetTime: u.FlowResetTime, Num: u.Num,
|
||||
CreatedTime: u.CreatedTime, Status: u.Status,
|
||||
}
|
||||
if quota := quotaMap[u.ID]; quota != nil {
|
||||
b.DailyQuotaGB = quota.DailyLimitGB
|
||||
b.MonthlyQuotaGB = quota.MonthlyLimitGB
|
||||
b.DisabledByQuota = quota.DisabledByQuota
|
||||
b.QuotaDisabledAt = quota.DisabledAt
|
||||
}
|
||||
if u.UpdatedTime.Valid {
|
||||
b.UpdatedTime = u.UpdatedTime.Int64
|
||||
}
|
||||
@@ -1882,14 +1894,6 @@ func (r *Repository) exportTunnels() ([]model.TunnelBackup, error) {
|
||||
if err := r.db.Order("inx ASC, id ASC").Find(&tunnels).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
quotaIDs := make([]int64, 0, len(tunnels))
|
||||
for _, t := range tunnels {
|
||||
quotaIDs = append(quotaIDs, t.ID)
|
||||
}
|
||||
quotaMap, err := r.ListTunnelQuotaViewsByTunnelIDs(quotaIDs, time.Now())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out := make([]model.TunnelBackup, 0, len(tunnels))
|
||||
for _, t := range tunnels {
|
||||
b := model.TunnelBackup{
|
||||
@@ -1898,12 +1902,6 @@ func (r *Repository) exportTunnels() ([]model.TunnelBackup, error) {
|
||||
CreatedTime: t.CreatedTime, UpdatedTime: t.UpdatedTime,
|
||||
Status: t.Status, Inx: t.Inx, IPPreference: t.IPPreference,
|
||||
}
|
||||
if quota := quotaMap[t.ID]; quota != nil {
|
||||
b.DailyQuotaGB = quota.DailyLimitGB
|
||||
b.MonthlyQuotaGB = quota.MonthlyLimitGB
|
||||
b.DisabledByQuota = quota.DisabledByQuota
|
||||
b.QuotaDisabledAt = quota.DisabledAt
|
||||
}
|
||||
if t.InIP.Valid {
|
||||
b.InIP = t.InIP.String
|
||||
}
|
||||
@@ -2200,6 +2198,39 @@ func importUsers(tx *gorm.DB, users []model.UserBackup, now int64) (int, error)
|
||||
if err != nil {
|
||||
return count, err
|
||||
}
|
||||
if u.DailyQuotaGB > 0 || u.MonthlyQuotaGB > 0 || u.DisabledByQuota != 0 || u.QuotaDisabledAt > 0 {
|
||||
current := time.UnixMilli(now)
|
||||
dayKey := int64(current.Year()*10000 + int(current.Month())*100 + current.Day())
|
||||
monthKey := int64(current.Year()*100 + int(current.Month()))
|
||||
quotaItem := model.UserQuota{
|
||||
UserID: u.ID,
|
||||
DailyLimitGB: u.DailyQuotaGB,
|
||||
MonthlyLimitGB: u.MonthlyQuotaGB,
|
||||
DailyUsedBytes: 0,
|
||||
MonthlyUsedBytes: 0,
|
||||
DayKey: dayKey,
|
||||
MonthKey: monthKey,
|
||||
DisabledByQuota: u.DisabledByQuota,
|
||||
DisabledAt: u.QuotaDisabledAt,
|
||||
PausedForwardIDs: "",
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}
|
||||
err = tx.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "user_id"}},
|
||||
DoUpdates: clause.AssignmentColumns([]string{
|
||||
"daily_limit_gb", "monthly_limit_gb", "daily_used_bytes", "monthly_used_bytes",
|
||||
"day_key", "month_key", "disabled_by_quota", "disabled_at", "paused_forward_ids", "updated_time",
|
||||
}),
|
||||
}).Create("aItem).Error
|
||||
if err != nil {
|
||||
return count, err
|
||||
}
|
||||
} else {
|
||||
if err := tx.Where("user_id = ?", u.ID).Delete(&model.UserQuota{}).Error; err != nil {
|
||||
return count, err
|
||||
}
|
||||
}
|
||||
count++
|
||||
}
|
||||
return count, nil
|
||||
@@ -2277,26 +2308,6 @@ func importTunnels(tx *gorm.DB, tunnels []model.TunnelBackup, now int64) (int, e
|
||||
if err != nil {
|
||||
return count, err
|
||||
}
|
||||
quotaItem := model.TunnelQuota{
|
||||
TunnelID: t.ID,
|
||||
DailyLimitGB: t.DailyQuotaGB,
|
||||
MonthlyLimitGB: t.MonthlyQuotaGB,
|
||||
DisabledByQuota: t.DisabledByQuota,
|
||||
DisabledAt: t.QuotaDisabledAt,
|
||||
DayKey: int64(time.Now().Year()*10000 + int(time.Now().Month())*100 + time.Now().Day()),
|
||||
MonthKey: int64(time.Now().Year()*100 + int(time.Now().Month())),
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}
|
||||
err = tx.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "tunnel_id"}},
|
||||
DoUpdates: clause.AssignmentColumns([]string{
|
||||
"daily_limit_gb", "monthly_limit_gb", "disabled_by_quota", "disabled_at", "updated_time",
|
||||
}),
|
||||
}).Create("aItem).Error
|
||||
if err != nil {
|
||||
return count, err
|
||||
}
|
||||
for _, ct := range t.ChainTunnels {
|
||||
chainItem := model.ChainTunnel{
|
||||
ID: ct.ID,
|
||||
|
||||
@@ -145,6 +145,9 @@ func (r *Repository) DeleteUserCascade(userID int64) error {
|
||||
if err := tx.Where("user_id = ?", userID).Delete(&model.StatisticsFlow{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Where("user_id = ?", userID).Delete(&model.UserQuota{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Where("id = ?", userID).Delete(&model.User{}).Error
|
||||
})
|
||||
}
|
||||
|
||||
+70
-83
@@ -13,21 +13,21 @@ import (
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
const tunnelQuotaBytesPerGB int64 = 1024 * 1024 * 1024
|
||||
const userQuotaBytesPerGB int64 = 1024 * 1024 * 1024
|
||||
|
||||
type TunnelQuotaRelease struct {
|
||||
TunnelID int64
|
||||
ForwardIDs []int64
|
||||
EnableTunnel bool
|
||||
type UserQuotaRelease struct {
|
||||
UserID int64
|
||||
ForwardIDs []int64
|
||||
UnblockUser bool
|
||||
}
|
||||
|
||||
func tunnelQuotaWindowKeys(now time.Time) (int64, int64) {
|
||||
func userQuotaWindowKeys(now time.Time) (int64, int64) {
|
||||
return int64(now.Year()*10000 + int(now.Month())*100 + now.Day()), int64(now.Year()*100 + int(now.Month()))
|
||||
}
|
||||
|
||||
func cloneTunnelQuotaView(q model.TunnelQuota) *model.TunnelQuotaView {
|
||||
return &model.TunnelQuotaView{
|
||||
TunnelID: q.TunnelID,
|
||||
func cloneUserQuotaView(q model.UserQuota) *model.UserQuotaView {
|
||||
return &model.UserQuotaView{
|
||||
UserID: q.UserID,
|
||||
DailyLimitGB: q.DailyLimitGB,
|
||||
MonthlyLimitGB: q.MonthlyLimitGB,
|
||||
DailyUsedBytes: q.DailyUsedBytes,
|
||||
@@ -40,11 +40,11 @@ func cloneTunnelQuotaView(q model.TunnelQuota) *model.TunnelQuotaView {
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeTunnelQuotaView(view *model.TunnelQuotaView, now time.Time) *model.TunnelQuotaView {
|
||||
func normalizeUserQuotaView(view *model.UserQuotaView, now time.Time) *model.UserQuotaView {
|
||||
if view == nil {
|
||||
return nil
|
||||
}
|
||||
dayKey, monthKey := tunnelQuotaWindowKeys(now)
|
||||
dayKey, monthKey := userQuotaWindowKeys(now)
|
||||
out := *view
|
||||
if out.DayKey != dayKey {
|
||||
out.DayKey = dayKey
|
||||
@@ -57,14 +57,14 @@ func normalizeTunnelQuotaView(view *model.TunnelQuotaView, now time.Time) *model
|
||||
return &out
|
||||
}
|
||||
|
||||
func tunnelQuotaExceeded(view *model.TunnelQuotaView) bool {
|
||||
func userQuotaExceeded(view *model.UserQuotaView) bool {
|
||||
if view == nil {
|
||||
return false
|
||||
}
|
||||
if view.DailyLimitGB > 0 && view.DailyUsedBytes >= view.DailyLimitGB*tunnelQuotaBytesPerGB {
|
||||
if view.DailyLimitGB > 0 && view.DailyUsedBytes >= view.DailyLimitGB*userQuotaBytesPerGB {
|
||||
return true
|
||||
}
|
||||
if view.MonthlyLimitGB > 0 && view.MonthlyUsedBytes >= view.MonthlyLimitGB*tunnelQuotaBytesPerGB {
|
||||
if view.MonthlyLimitGB > 0 && view.MonthlyUsedBytes >= view.MonthlyLimitGB*userQuotaBytesPerGB {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
@@ -107,13 +107,13 @@ func joinPausedForwardIDs(ids []int64) string {
|
||||
return strings.Join(parts, ",")
|
||||
}
|
||||
|
||||
func (r *Repository) loadOrCreateTunnelQuotaTx(tx *gorm.DB, tunnelID int64, now time.Time) (*model.TunnelQuota, error) {
|
||||
func (r *Repository) loadOrCreateUserQuotaTx(tx *gorm.DB, userID int64, now time.Time) (*model.UserQuota, error) {
|
||||
if tx == nil {
|
||||
return nil, errors.New("database unavailable")
|
||||
}
|
||||
dayKey, monthKey := tunnelQuotaWindowKeys(now)
|
||||
q := &model.TunnelQuota{}
|
||||
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Where("tunnel_id = ?", tunnelID).First(q).Error
|
||||
dayKey, monthKey := userQuotaWindowKeys(now)
|
||||
q := &model.UserQuota{}
|
||||
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Where("user_id = ?", userID).First(q).Error
|
||||
if err == nil {
|
||||
return q, nil
|
||||
}
|
||||
@@ -121,8 +121,8 @@ func (r *Repository) loadOrCreateTunnelQuotaTx(tx *gorm.DB, tunnelID int64, now
|
||||
return nil, err
|
||||
}
|
||||
nowMs := now.UnixMilli()
|
||||
q = &model.TunnelQuota{
|
||||
TunnelID: tunnelID,
|
||||
q = &model.UserQuota{
|
||||
UserID: userID,
|
||||
DayKey: dayKey,
|
||||
MonthKey: monthKey,
|
||||
CreatedTime: nowMs,
|
||||
@@ -135,12 +135,12 @@ func (r *Repository) loadOrCreateTunnelQuotaTx(tx *gorm.DB, tunnelID int64, now
|
||||
return q, nil
|
||||
}
|
||||
|
||||
func applyTunnelQuotaWindowRoll(q *model.TunnelQuota, now time.Time) bool {
|
||||
func applyUserQuotaWindowRoll(q *model.UserQuota, now time.Time) bool {
|
||||
if q == nil {
|
||||
return false
|
||||
}
|
||||
changed := false
|
||||
dayKey, monthKey := tunnelQuotaWindowKeys(now)
|
||||
dayKey, monthKey := userQuotaWindowKeys(now)
|
||||
if q.DayKey != dayKey {
|
||||
q.DayKey = dayKey
|
||||
q.DailyUsedBytes = 0
|
||||
@@ -154,18 +154,18 @@ func applyTunnelQuotaWindowRoll(q *model.TunnelQuota, now time.Time) bool {
|
||||
return changed
|
||||
}
|
||||
|
||||
func (r *Repository) SaveTunnelQuotaConfigTx(tx *gorm.DB, tunnelID, dailyLimitGB, monthlyLimitGB int64, now int64) error {
|
||||
func (r *Repository) SaveUserQuotaConfigTx(tx *gorm.DB, userID, dailyLimitGB, monthlyLimitGB int64, now int64) error {
|
||||
if tx == nil {
|
||||
return errors.New("database unavailable")
|
||||
}
|
||||
if tunnelID <= 0 {
|
||||
return errors.New("tunnel id is required")
|
||||
if userID <= 0 {
|
||||
return errors.New("user id is required")
|
||||
}
|
||||
if dailyLimitGB < 0 || monthlyLimitGB < 0 {
|
||||
return errors.New("quota limit cannot be negative")
|
||||
}
|
||||
current := time.UnixMilli(now)
|
||||
q, err := r.loadOrCreateTunnelQuotaTx(tx, tunnelID, current)
|
||||
q, err := r.loadOrCreateUserQuotaTx(tx, userID, current)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -175,69 +175,69 @@ func (r *Repository) SaveTunnelQuotaConfigTx(tx *gorm.DB, tunnelID, dailyLimitGB
|
||||
"updated_time": now,
|
||||
}
|
||||
if q.DayKey == 0 || q.MonthKey == 0 {
|
||||
dayKey, monthKey := tunnelQuotaWindowKeys(current)
|
||||
dayKey, monthKey := userQuotaWindowKeys(current)
|
||||
updates["day_key"] = dayKey
|
||||
updates["month_key"] = monthKey
|
||||
}
|
||||
return tx.Model(&model.TunnelQuota{}).Where("tunnel_id = ?", tunnelID).Updates(updates).Error
|
||||
return tx.Model(&model.UserQuota{}).Where("user_id = ?", userID).Updates(updates).Error
|
||||
}
|
||||
|
||||
func (r *Repository) ListTunnelQuotaViewsByTunnelIDs(tunnelIDs []int64, now time.Time) (map[int64]*model.TunnelQuotaView, error) {
|
||||
func (r *Repository) ListUserQuotaViewsByUserIDs(userIDs []int64, now time.Time) (map[int64]*model.UserQuotaView, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
out := make(map[int64]*model.TunnelQuotaView)
|
||||
if len(tunnelIDs) == 0 {
|
||||
out := make(map[int64]*model.UserQuotaView)
|
||||
if len(userIDs) == 0 {
|
||||
return out, nil
|
||||
}
|
||||
var rows []model.TunnelQuota
|
||||
if err := r.db.Where("tunnel_id IN ?", tunnelIDs).Find(&rows).Error; err != nil {
|
||||
var rows []model.UserQuota
|
||||
if err := r.db.Where("user_id IN ?", userIDs).Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, row := range rows {
|
||||
out[row.TunnelID] = normalizeTunnelQuotaView(cloneTunnelQuotaView(row), now)
|
||||
out[row.UserID] = normalizeUserQuotaView(cloneUserQuotaView(row), now)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (r *Repository) GetTunnelQuotaView(tunnelID int64, now time.Time) (*model.TunnelQuotaView, error) {
|
||||
func (r *Repository) GetUserQuotaView(userID int64, now time.Time) (*model.UserQuotaView, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
if tunnelID <= 0 {
|
||||
if userID <= 0 {
|
||||
return nil, nil
|
||||
}
|
||||
var row model.TunnelQuota
|
||||
err := r.db.Where("tunnel_id = ?", tunnelID).First(&row).Error
|
||||
var row model.UserQuota
|
||||
err := r.db.Where("user_id = ?", userID).First(&row).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, nil
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return normalizeTunnelQuotaView(cloneTunnelQuotaView(row), now), nil
|
||||
return normalizeUserQuotaView(cloneUserQuotaView(row), now), nil
|
||||
}
|
||||
|
||||
func (r *Repository) AddTunnelQuotaUsage(tunnelID int64, usedBytes int64, now time.Time) (*model.TunnelQuotaView, error) {
|
||||
func (r *Repository) AddUserQuotaUsage(userID int64, usedBytes int64, now time.Time) (*model.UserQuotaView, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
if tunnelID <= 0 {
|
||||
if userID <= 0 {
|
||||
return nil, nil
|
||||
}
|
||||
result := &model.TunnelQuotaView{}
|
||||
result := &model.UserQuotaView{}
|
||||
err := r.db.Transaction(func(tx *gorm.DB) error {
|
||||
q, err := r.loadOrCreateTunnelQuotaTx(tx, tunnelID, now)
|
||||
q, err := r.loadOrCreateUserQuotaTx(tx, userID, now)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
applyTunnelQuotaWindowRoll(q, now)
|
||||
applyUserQuotaWindowRoll(q, now)
|
||||
if usedBytes > 0 {
|
||||
q.DailyUsedBytes += usedBytes
|
||||
q.MonthlyUsedBytes += usedBytes
|
||||
}
|
||||
q.UpdatedTime = now.UnixMilli()
|
||||
if err := tx.Model(&model.TunnelQuota{}).Where("tunnel_id = ?", tunnelID).Updates(map[string]interface{}{
|
||||
if err := tx.Model(&model.UserQuota{}).Where("user_id = ?", userID).Updates(map[string]interface{}{
|
||||
"daily_used_bytes": q.DailyUsedBytes,
|
||||
"monthly_used_bytes": q.MonthlyUsedBytes,
|
||||
"day_key": q.DayKey,
|
||||
@@ -246,23 +246,23 @@ func (r *Repository) AddTunnelQuotaUsage(tunnelID int64, usedBytes int64, now ti
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
*result = *cloneTunnelQuotaView(*q)
|
||||
*result = *cloneUserQuotaView(*q)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return normalizeTunnelQuotaView(result, now), nil
|
||||
return normalizeUserQuotaView(result, now), nil
|
||||
}
|
||||
|
||||
func (r *Repository) MarkTunnelQuotaDisabled(tunnelID int64, pausedForwardIDs []int64, now int64) error {
|
||||
func (r *Repository) MarkUserQuotaDisabled(userID int64, pausedForwardIDs []int64, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
if tunnelID <= 0 {
|
||||
return errors.New("tunnel id is required")
|
||||
if userID <= 0 {
|
||||
return errors.New("user id is required")
|
||||
}
|
||||
return r.db.Model(&model.TunnelQuota{}).Where("tunnel_id = ?", tunnelID).Updates(map[string]interface{}{
|
||||
return r.db.Model(&model.UserQuota{}).Where("user_id = ?", userID).Updates(map[string]interface{}{
|
||||
"disabled_by_quota": 1,
|
||||
"disabled_at": now,
|
||||
"paused_forward_ids": joinPausedForwardIDs(pausedForwardIDs),
|
||||
@@ -270,12 +270,12 @@ func (r *Repository) MarkTunnelQuotaDisabled(tunnelID int64, pausedForwardIDs []
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) ResetTunnelQuotaUsage(tunnelID int64, scope string, now time.Time) (*TunnelQuotaRelease, error) {
|
||||
func (r *Repository) ResetUserQuotaUsage(userID int64, scope string, now time.Time) (*UserQuotaRelease, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
if tunnelID <= 0 {
|
||||
return nil, errors.New("tunnel id is required")
|
||||
if userID <= 0 {
|
||||
return nil, errors.New("user id is required")
|
||||
}
|
||||
scope = strings.TrimSpace(strings.ToLower(scope))
|
||||
if scope == "" {
|
||||
@@ -284,13 +284,13 @@ func (r *Repository) ResetTunnelQuotaUsage(tunnelID int64, scope string, now tim
|
||||
if scope != "daily" && scope != "monthly" && scope != "all" {
|
||||
return nil, fmt.Errorf("unsupported quota reset scope: %s", scope)
|
||||
}
|
||||
var release *TunnelQuotaRelease
|
||||
var release *UserQuotaRelease
|
||||
err := r.db.Transaction(func(tx *gorm.DB) error {
|
||||
q, err := r.loadOrCreateTunnelQuotaTx(tx, tunnelID, now)
|
||||
q, err := r.loadOrCreateUserQuotaTx(tx, userID, now)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
applyTunnelQuotaWindowRoll(q, now)
|
||||
applyUserQuotaWindowRoll(q, now)
|
||||
switch scope {
|
||||
case "daily":
|
||||
q.DailyUsedBytes = 0
|
||||
@@ -301,15 +301,15 @@ func (r *Repository) ResetTunnelQuotaUsage(tunnelID int64, scope string, now tim
|
||||
q.MonthlyUsedBytes = 0
|
||||
}
|
||||
q.UpdatedTime = now.UnixMilli()
|
||||
release = &TunnelQuotaRelease{TunnelID: tunnelID}
|
||||
if q.DisabledByQuota == 1 && !tunnelQuotaExceeded(cloneTunnelQuotaView(*q)) {
|
||||
release.EnableTunnel = true
|
||||
release = &UserQuotaRelease{UserID: userID}
|
||||
if q.DisabledByQuota == 1 && !userQuotaExceeded(cloneUserQuotaView(*q)) {
|
||||
release.UnblockUser = true
|
||||
release.ForwardIDs = parsePausedForwardIDs(q.PausedForwardIDs)
|
||||
q.DisabledByQuota = 0
|
||||
q.DisabledAt = 0
|
||||
q.PausedForwardIDs = ""
|
||||
}
|
||||
return tx.Model(&model.TunnelQuota{}).Where("tunnel_id = ?", tunnelID).Updates(map[string]interface{}{
|
||||
return tx.Model(&model.UserQuota{}).Where("user_id = ?", userID).Updates(map[string]interface{}{
|
||||
"daily_used_bytes": q.DailyUsedBytes,
|
||||
"monthly_used_bytes": q.MonthlyUsedBytes,
|
||||
"day_key": q.DayKey,
|
||||
@@ -326,23 +326,23 @@ func (r *Repository) ResetTunnelQuotaUsage(tunnelID int64, scope string, now tim
|
||||
return release, nil
|
||||
}
|
||||
|
||||
func (r *Repository) RollTunnelQuotaWindows(now time.Time) ([]TunnelQuotaRelease, error) {
|
||||
func (r *Repository) RollUserQuotaWindows(now time.Time) ([]UserQuotaRelease, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var releases []TunnelQuotaRelease
|
||||
var releases []UserQuotaRelease
|
||||
err := r.db.Transaction(func(tx *gorm.DB) error {
|
||||
var rows []model.TunnelQuota
|
||||
var rows []model.UserQuota
|
||||
if err := tx.Find(&rows).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
nowMs := now.UnixMilli()
|
||||
for _, row := range rows {
|
||||
q := row
|
||||
changed := applyTunnelQuotaWindowRoll(&q, now)
|
||||
release := TunnelQuotaRelease{TunnelID: q.TunnelID}
|
||||
if q.DisabledByQuota == 1 && !tunnelQuotaExceeded(cloneTunnelQuotaView(q)) {
|
||||
release.EnableTunnel = true
|
||||
changed := applyUserQuotaWindowRoll(&q, now)
|
||||
release := UserQuotaRelease{UserID: q.UserID}
|
||||
if q.DisabledByQuota == 1 && !userQuotaExceeded(cloneUserQuotaView(q)) {
|
||||
release.UnblockUser = true
|
||||
release.ForwardIDs = parsePausedForwardIDs(q.PausedForwardIDs)
|
||||
q.DisabledByQuota = 0
|
||||
q.DisabledAt = 0
|
||||
@@ -353,7 +353,7 @@ func (r *Repository) RollTunnelQuotaWindows(now time.Time) ([]TunnelQuotaRelease
|
||||
continue
|
||||
}
|
||||
q.UpdatedTime = nowMs
|
||||
if err := tx.Model(&model.TunnelQuota{}).Where("tunnel_id = ?", q.TunnelID).Updates(map[string]interface{}{
|
||||
if err := tx.Model(&model.UserQuota{}).Where("user_id = ?", q.UserID).Updates(map[string]interface{}{
|
||||
"daily_used_bytes": q.DailyUsedBytes,
|
||||
"monthly_used_bytes": q.MonthlyUsedBytes,
|
||||
"day_key": q.DayKey,
|
||||
@@ -365,7 +365,7 @@ func (r *Repository) RollTunnelQuotaWindows(now time.Time) ([]TunnelQuotaRelease
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if release.EnableTunnel {
|
||||
if release.UnblockUser {
|
||||
releases = append(releases, release)
|
||||
}
|
||||
}
|
||||
@@ -376,16 +376,3 @@ func (r *Repository) RollTunnelQuotaWindows(now time.Time) ([]TunnelQuotaRelease
|
||||
}
|
||||
return releases, nil
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateTunnelStatus(tunnelID int64, status int, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
if tunnelID <= 0 {
|
||||
return errors.New("tunnel id is required")
|
||||
}
|
||||
return r.db.Model(&model.Tunnel{}).Where("id = ?", tunnelID).Updates(map[string]interface{}{
|
||||
"status": status,
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
@@ -0,0 +1,202 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestForwardBatchDeleteReturnsFailureReasonsContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, _ := setupContractRouter(t, secret)
|
||||
adminToken := mustAdminToken(t, secret)
|
||||
|
||||
out := postBatchRequest(t, router, adminToken, "/api/v1/forward/batch-delete", `{"ids":[999]}`)
|
||||
result := mustBatchResult(t, out)
|
||||
assertBatchFailureReasonContains(t, result, "转发不存在")
|
||||
}
|
||||
|
||||
func TestForwardBatchPauseReturnsFailureReasonsContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, _ := setupContractRouter(t, secret)
|
||||
adminToken := mustAdminToken(t, secret)
|
||||
|
||||
out := postBatchRequest(t, router, adminToken, "/api/v1/forward/batch-pause", `{"ids":[999]}`)
|
||||
result := mustBatchResult(t, out)
|
||||
assertBatchFailureReasonContains(t, result, "转发不存在")
|
||||
}
|
||||
|
||||
func TestForwardBatchResumeReturnsFailureReasonsContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
adminToken := mustAdminToken(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
forwardID := seedForwardForBatchAction(t, repo, batchForwardSeedOptions{
|
||||
Now: now,
|
||||
TunnelName: "resume-detail-tunnel",
|
||||
ForwardName: "resume-detail-forward",
|
||||
CreateUserTunnel: true,
|
||||
UserTunnelStatus: 0,
|
||||
})
|
||||
|
||||
out := postBatchRequest(t, router, adminToken, "/api/v1/forward/batch-resume", `{"ids":[`+jsonNumber(forwardID)+`]}`)
|
||||
result := mustBatchResult(t, out)
|
||||
assertBatchFailureNameAndReason(t, result, "resume-detail-forward", "该隧道已禁用")
|
||||
}
|
||||
|
||||
func TestForwardBatchChangeTunnelReturnsFailureReasonsContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
adminToken := mustAdminToken(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
forwardID := seedForwardForBatchAction(t, repo, batchForwardSeedOptions{
|
||||
Now: now,
|
||||
TunnelName: "change-detail-tunnel",
|
||||
ForwardName: "change-detail-forward",
|
||||
})
|
||||
tunnelID := mustQueryInt64(t, repo, `SELECT tunnel_id FROM forward WHERE id = ?`, forwardID)
|
||||
|
||||
payload := `{"forwardIds":[` + jsonNumber(forwardID) + `],"targetTunnelId":` + jsonNumber(tunnelID) + `}`
|
||||
out := postBatchRequest(t, router, adminToken, "/api/v1/forward/batch-change-tunnel", payload)
|
||||
result := mustBatchResult(t, out)
|
||||
assertBatchFailureNameAndReason(t, result, "change-detail-forward", "规则已在目标隧道中")
|
||||
}
|
||||
|
||||
func TestTunnelBatchDeleteReturnsFailureReasonsContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, _ := setupContractRouter(t, secret)
|
||||
adminToken := mustAdminToken(t, secret)
|
||||
|
||||
out := postBatchRequest(t, router, adminToken, "/api/v1/tunnel/batch-delete", `{"ids":[999]}`)
|
||||
result := mustBatchResult(t, out)
|
||||
assertBatchFailureReasonContains(t, result, "隧道不存在")
|
||||
}
|
||||
|
||||
type batchForwardSeedOptions struct {
|
||||
Now int64
|
||||
TunnelName string
|
||||
ForwardName string
|
||||
CreateUserTunnel bool
|
||||
UserTunnelStatus int
|
||||
}
|
||||
|
||||
func mustAdminToken(t *testing.T, secret string) string {
|
||||
t.Helper()
|
||||
token, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
return token
|
||||
}
|
||||
|
||||
func postBatchRequest(t *testing.T, router http.Handler, token, path, payload string) response.R {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(http.MethodPost, path, bytes.NewBufferString(payload))
|
||||
req.Header.Set("Authorization", token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
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 API success envelope, got code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func mustBatchResult(t *testing.T, out response.R) map[string]interface{} {
|
||||
t.Helper()
|
||||
result, ok := out.Data.(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected map result, got %T", out.Data)
|
||||
}
|
||||
if int(result["failCount"].(float64)) != 1 {
|
||||
t.Fatalf("expected failCount=1, got %v", result["failCount"])
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func assertBatchFailureReasonContains(t *testing.T, result map[string]interface{}, snippet string) {
|
||||
t.Helper()
|
||||
failures, ok := result["failures"].([]interface{})
|
||||
if !ok || len(failures) != 1 {
|
||||
t.Fatalf("expected exactly one failure detail, got %#v", result["failures"])
|
||||
}
|
||||
first, ok := failures[0].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected failure detail object, got %T", failures[0])
|
||||
}
|
||||
reason, _ := first["reason"].(string)
|
||||
if !strings.Contains(reason, snippet) {
|
||||
t.Fatalf("expected failure reason to contain %q, got %q", snippet, reason)
|
||||
}
|
||||
}
|
||||
|
||||
func assertBatchFailureNameAndReason(t *testing.T, result map[string]interface{}, expectedName, reasonSnippet string) {
|
||||
t.Helper()
|
||||
failures, ok := result["failures"].([]interface{})
|
||||
if !ok || len(failures) != 1 {
|
||||
t.Fatalf("expected exactly one failure detail, got %#v", result["failures"])
|
||||
}
|
||||
first, ok := failures[0].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected failure detail object, got %T", failures[0])
|
||||
}
|
||||
gotName, _ := first["name"].(string)
|
||||
if strings.TrimSpace(gotName) != expectedName {
|
||||
t.Fatalf("expected failure name %q, got %q", expectedName, gotName)
|
||||
}
|
||||
reason, _ := first["reason"].(string)
|
||||
if !strings.Contains(reason, reasonSnippet) {
|
||||
t.Fatalf("expected failure reason to contain %q, got %q", reasonSnippet, reason)
|
||||
}
|
||||
}
|
||||
|
||||
func seedForwardForBatchAction(t *testing.T, repo *repo.Repository, opts batchForwardSeedOptions) int64 {
|
||||
t.Helper()
|
||||
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(2, 'batch_action_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, opts.Now, opts.Now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, opts.TunnelName, 1.0, 1, "tls", 99999, opts.Now, opts.Now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, repo, opts.TunnelName)
|
||||
|
||||
if opts.CreateUserTunnel {
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(20, 2, ?, NULL, 999, 99999, 0, 0, 1, 2727251700000, ?)
|
||||
`, tunnelID, opts.UserTunnelStatus).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||
VALUES(2, 'batch_action_user', ?, ?, '1.1.1.1:443', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
`, opts.ForwardName, tunnelID, opts.Now, opts.Now).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
|
||||
return mustLastInsertID(t, repo, opts.ForwardName)
|
||||
}
|
||||
@@ -0,0 +1,161 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
func TestForwardBatchRedeployReturnsFailureReasonsContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
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(2, 'batch_redeploy_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "batch-redeploy-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, repo, "batch-redeploy-tunnel")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, 0)
|
||||
`, 2, "batch_redeploy_user", "redeploy-forward", tunnelID, "1.1.1.1:443", "fifo", now, now).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
forwardID := mustLastInsertID(t, repo, "redeploy-forward")
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/batch-redeploy", bytes.NewBufferString(`{"ids":[`+jsonNumber(forwardID)+`]}`))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
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 API success envelope, got code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
result, ok := out.Data.(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected map result, got %T", out.Data)
|
||||
}
|
||||
if int(result["failCount"].(float64)) != 1 {
|
||||
t.Fatalf("expected failCount=1, got %v", result["failCount"])
|
||||
}
|
||||
if int(result["successCount"].(float64)) != 0 {
|
||||
t.Fatalf("expected successCount=0, got %v", result["successCount"])
|
||||
}
|
||||
|
||||
failures, ok := result["failures"].([]interface{})
|
||||
if !ok || len(failures) != 1 {
|
||||
t.Fatalf("expected exactly one failure detail, got %#v", result["failures"])
|
||||
}
|
||||
first, ok := failures[0].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected failure detail object, got %T", failures[0])
|
||||
}
|
||||
if gotName := strings.TrimSpace(first["name"].(string)); gotName != "redeploy-forward" {
|
||||
t.Fatalf("expected failure name redeploy-forward, got %q", gotName)
|
||||
}
|
||||
reason, _ := first["reason"].(string)
|
||||
if !strings.Contains(reason, "转发入口端口不存在") {
|
||||
t.Fatalf("expected forward failure reason to mention missing entry port, got %q", reason)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelBatchRedeployReturnsFailureReasonsContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "broken-redeploy-tunnel", 1.0, 2, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, repo, "broken-redeploy-tunnel")
|
||||
|
||||
if 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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "entry-only-node", "entry-only-secret", "10.0.0.20", "10.0.0.20", "", "20000-20010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node: %v", err)
|
||||
}
|
||||
entryNodeID := mustLastInsertID(t, repo, "entry-only-node")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 1, ?, 20001, 'round', 1, 'tls')
|
||||
`, tunnelID, entryNodeID).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/batch-redeploy", bytes.NewBufferString(`{"ids":[`+jsonNumber(tunnelID)+`]}`))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
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 API success envelope, got code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
result, ok := out.Data.(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected map result, got %T", out.Data)
|
||||
}
|
||||
if int(result["failCount"].(float64)) != 1 {
|
||||
t.Fatalf("expected failCount=1, got %v", result["failCount"])
|
||||
}
|
||||
|
||||
failures, ok := result["failures"].([]interface{})
|
||||
if !ok || len(failures) != 1 {
|
||||
t.Fatalf("expected exactly one failure detail, got %#v", result["failures"])
|
||||
}
|
||||
first, ok := failures[0].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected failure detail object, got %T", failures[0])
|
||||
}
|
||||
if gotName := strings.TrimSpace(first["name"].(string)); gotName != "broken-redeploy-tunnel" {
|
||||
t.Fatalf("expected failure name broken-redeploy-tunnel, got %q", gotName)
|
||||
}
|
||||
reason, _ := first["reason"].(string)
|
||||
if !strings.Contains(reason, "转发链目标不能为空") {
|
||||
t.Fatalf("expected tunnel failure reason to mention missing target, got %q", reason)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,187 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
func TestIssue313_EntryPortCrossTunnelConflictContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
insertNode := func(name, ip, portRange string) int64 {
|
||||
if 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, "", portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node %s: %v", name, err)
|
||||
}
|
||||
return mustLastInsertID(t, repo, name)
|
||||
}
|
||||
|
||||
entryA := insertNode("issue313-entry-a", "10.100.0.1", "2000-2010")
|
||||
entryB1 := insertNode("issue313-entry-b1", "10.100.0.2", "2000-2010")
|
||||
entryB2 := insertNode("issue313-entry-b2", "10.100.0.3", "2000-2010")
|
||||
chainA := insertNode("issue313-chain-a", "10.100.0.4", "3000-3010")
|
||||
chainB := insertNode("issue313-chain-b", "10.100.0.5", "3000-3010")
|
||||
exitA := insertNode("issue313-exit-a", "10.100.0.6", "4000-4010")
|
||||
exitB := insertNode("issue313-exit-b", "10.100.0.7", "4000-4010")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "issue313-tunnel-a", 1.0, 2, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel a: %v", err)
|
||||
}
|
||||
tunnelAID := mustLastInsertID(t, repo, "issue313-tunnel-a")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 1, ?, 2000, 'round', 1, 'tls')
|
||||
`, tunnelAID, entryA).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel entry a: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 2, ?, 3000, 'round', 1, 'tls')
|
||||
`, tunnelAID, chainA).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel chain a: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 3, ?, 4000, 'round', 1, 'tls')
|
||||
`, tunnelAID, exitA).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel exit a: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "issue313-tunnel-b", 1.0, 2, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel b: %v", err)
|
||||
}
|
||||
tunnelBID := mustLastInsertID(t, repo, "issue313-tunnel-b")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 1, ?, 2000, 'round', 1, 'tls')
|
||||
`, tunnelBID, entryB1).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel entry b1: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 2, ?, 3000, 'round', 1, 'tls')
|
||||
`, tunnelBID, chainB).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel chain b: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 3, ?, 4000, 'round', 1, 'tls')
|
||||
`, tunnelBID, exitB).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel exit b: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(3131, 1, ?, NULL, 999, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`, tunnelAID).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel for tunnel a: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||
VALUES(1, 'admin_user', 'issue313-forward-a', ?, '1.1.1.1:443', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
`, tunnelAID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert forward a: %v", err)
|
||||
}
|
||||
forwardAID := mustLastInsertID(t, repo, "issue313-forward-a")
|
||||
|
||||
if err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardAID, entryA, 2000).Error; err != nil {
|
||||
t.Fatalf("insert forward_port a: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(3132, 1, ?, NULL, 999, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`, tunnelBID).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel for tunnel b: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||
VALUES(1, 'admin_user', 'issue313-forward-b', ?, '2.2.2.2:443', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
`, tunnelBID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert forward b: %v", err)
|
||||
}
|
||||
forwardBID := mustLastInsertID(t, repo, "issue313-forward-b")
|
||||
|
||||
if err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardBID, entryB1, 2000).Error; err != nil {
|
||||
t.Fatalf("insert forward_port b: %v", err)
|
||||
}
|
||||
|
||||
payload := map[string]interface{}{
|
||||
"id": tunnelBID,
|
||||
"name": "issue313-tunnel-b",
|
||||
"type": 2,
|
||||
"flow": 99999,
|
||||
"trafficRatio": 1.0,
|
||||
"status": 1,
|
||||
"inNodeId": []map[string]interface{}{
|
||||
{"nodeId": entryB1, "protocol": "tls", "strategy": "round"},
|
||||
{"nodeId": entryB2, "protocol": "tls", "strategy": "round"},
|
||||
},
|
||||
"chainNodes": []interface{}{
|
||||
[]map[string]interface{}{{"nodeId": chainB, "protocol": "tls", "strategy": "round"}},
|
||||
},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": exitB, "protocol": "tls", "strategy": "round"},
|
||||
},
|
||||
}
|
||||
body, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal payload: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/update", bytes.NewReader(body))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
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 update failure due to cross-tunnel port conflict, got success with code 0")
|
||||
}
|
||||
|
||||
if !bytes.Contains(res.Body.Bytes(), []byte("端口")) && !bytes.Contains(res.Body.Bytes(), []byte("占用")) {
|
||||
t.Fatalf("expected port conflict error message, got %q", out.Msg)
|
||||
}
|
||||
|
||||
countB2 := mustQueryInt(t, repo, `SELECT COUNT(1) FROM forward_port WHERE forward_id = ? AND node_id = ?`, forwardBID, entryB2)
|
||||
if countB2 > 0 {
|
||||
t.Fatalf("expected no forward_port record for entryB2, but found %d", countB2)
|
||||
}
|
||||
|
||||
chainCountB2 := mustQueryInt(t, repo, `SELECT COUNT(1) FROM chain_tunnel WHERE tunnel_id = ? AND node_id = ?`, tunnelBID, entryB2)
|
||||
if chainCountB2 > 0 {
|
||||
t.Fatalf("expected no chain_tunnel record for entryB2, but found %d", chainCountB2)
|
||||
}
|
||||
}
|
||||
+42
-37
@@ -13,21 +13,24 @@ import (
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
func TestForwardCreateBlockedWhenTunnelQuotaExceeded(t *testing.T) {
|
||||
func TestForwardCreateBlockedWhenUserQuotaExceeded(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
now := time.Now()
|
||||
nowMs := now.UnixMilli()
|
||||
dayKey := int64(now.Year()*10000 + int(now.Month())*100 + now.Day())
|
||||
monthKey := int64(now.Year()*100 + int(now.Month()))
|
||||
|
||||
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(2, 'quota_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, now, now).Error; err != nil {
|
||||
`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(1, 'quota_tunnel', 1.0, 1, 'tls', 1, ?, ?, 0, NULL, 0)
|
||||
`, now, now).Error; err != nil {
|
||||
VALUES(1, 'quota_tunnel', 1.0, 1, 'tls', 1, ?, ?, 1, NULL, 0)
|
||||
`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
@@ -37,10 +40,10 @@ func TestForwardCreateBlockedWhenTunnelQuotaExceeded(t *testing.T) {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel_quota(tunnel_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time)
|
||||
VALUES(1, 10, 0, ?, ?, 20260311, 202603, 1, ?, '', ?, ?)
|
||||
`, 11*contractBytesPerGB, 11*contractBytesPerGB, now, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel_quota: %v", err)
|
||||
INSERT INTO user_quota(user_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time)
|
||||
VALUES(2, 10, 0, ?, ?, ?, ?, 1, ?, '', ?, ?)
|
||||
`, 11*contractBytesPerGB, 11*contractBytesPerGB, dayKey, monthKey, nowMs, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user_quota: %v", err)
|
||||
}
|
||||
|
||||
token, err := auth.GenerateToken(2, "quota_user", 1, secret)
|
||||
@@ -59,28 +62,31 @@ func TestForwardCreateBlockedWhenTunnelQuotaExceeded(t *testing.T) {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code == 0 {
|
||||
t.Fatalf("expected non-zero code when tunnel quota exceeded")
|
||||
t.Fatalf("expected non-zero code when user quota exceeded")
|
||||
}
|
||||
if !strings.Contains(out.Msg, "配额") {
|
||||
t.Fatalf("expected quota error, got %q", out.Msg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardResumeBlockedWhenTunnelQuotaExceeded(t *testing.T) {
|
||||
func TestForwardResumeBlockedWhenUserQuotaExceeded(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
now := time.Now()
|
||||
nowMs := now.UnixMilli()
|
||||
dayKey := int64(now.Year()*10000 + int(now.Month())*100 + now.Day())
|
||||
monthKey := int64(now.Year()*100 + int(now.Month()))
|
||||
|
||||
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(2, 'quota_resume_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, now, now).Error; err != nil {
|
||||
`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(1, 'quota_resume_tunnel', 1.0, 1, 'tls', 1, ?, ?, 0, NULL, 0)
|
||||
`, now, now).Error; err != nil {
|
||||
VALUES(1, 'quota_resume_tunnel', 1.0, 1, 'tls', 1, ?, ?, 1, NULL, 0)
|
||||
`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
@@ -92,14 +98,14 @@ func TestForwardResumeBlockedWhenTunnelQuotaExceeded(t *testing.T) {
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||
VALUES(1, 2, 'quota_resume_user', 'quota_resume_forward', 1, '1.1.1.1:53', 'fifo', 0, 0, ?, ?, 0, 0)
|
||||
`, now, now).Error; err != nil {
|
||||
`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel_quota(tunnel_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time)
|
||||
VALUES(1, 10, 0, ?, ?, 20260311, 202603, 1, ?, '1', ?, ?)
|
||||
`, 11*contractBytesPerGB, 11*contractBytesPerGB, now, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel_quota: %v", err)
|
||||
INSERT INTO user_quota(user_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time)
|
||||
VALUES(2, 10, 0, ?, ?, ?, ?, 1, ?, '1', ?, ?)
|
||||
`, 11*contractBytesPerGB, 11*contractBytesPerGB, dayKey, monthKey, nowMs, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user_quota: %v", err)
|
||||
}
|
||||
|
||||
token, err := auth.GenerateToken(2, "quota_resume_user", 1, secret)
|
||||
@@ -118,7 +124,7 @@ func TestForwardResumeBlockedWhenTunnelQuotaExceeded(t *testing.T) {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code == 0 {
|
||||
t.Fatalf("expected non-zero code when tunnel quota exceeded")
|
||||
t.Fatalf("expected non-zero code when user quota exceeded")
|
||||
}
|
||||
if !strings.Contains(out.Msg, "配额") {
|
||||
t.Fatalf("expected quota error, got %q", out.Msg)
|
||||
@@ -129,29 +135,32 @@ func TestForwardResumeBlockedWhenTunnelQuotaExceeded(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelQuotaResetReEnablesTunnel(t *testing.T) {
|
||||
func TestUserQuotaResetClearsDisableFlag(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
now := time.Now()
|
||||
nowMs := now.UnixMilli()
|
||||
dayKey := int64(now.Year()*10000 + int(now.Month())*100 + now.Day())
|
||||
monthKey := int64(now.Year()*100 + int(now.Month()))
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(1, 'quota_reset_tunnel', 1.0, 1, 'tls', 1, ?, ?, 0, NULL, 0)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
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(2, 'quota_reset_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel_quota(tunnel_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time)
|
||||
VALUES(1, 10, 0, ?, ?, 20260311, 202603, 1, ?, '', ?, ?)
|
||||
`, 11*contractBytesPerGB, 11*contractBytesPerGB, now, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel_quota: %v", err)
|
||||
INSERT INTO user_quota(user_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time)
|
||||
VALUES(2, 10, 0, ?, ?, ?, ?, 1, ?, '', ?, ?)
|
||||
`, 11*contractBytesPerGB, 11*contractBytesPerGB, dayKey, monthKey, nowMs, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user_quota: %v", err)
|
||||
}
|
||||
|
||||
token, err := auth.GenerateToken(1, "admin", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/quota/reset", bytes.NewBufferString(`{"tunnelId":1,"scope":"all"}`))
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/quota/reset", bytes.NewBufferString(`{"userId":2,"scope":"all"}`))
|
||||
req.Header.Set("Authorization", token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
@@ -165,11 +174,7 @@ func TestTunnelQuotaResetReEnablesTunnel(t *testing.T) {
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected reset success, got code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
tunnelStatus := mustQueryInt(t, repo, `SELECT status FROM tunnel WHERE id = 1`)
|
||||
if tunnelStatus != 1 {
|
||||
t.Fatalf("expected tunnel to be re-enabled, got %d", tunnelStatus)
|
||||
}
|
||||
quotaDisabled := mustQueryInt(t, repo, `SELECT disabled_by_quota FROM tunnel_quota WHERE tunnel_id = 1`)
|
||||
quotaDisabled := mustQueryInt(t, repo, `SELECT disabled_by_quota FROM user_quota WHERE user_id = 2`)
|
||||
if quotaDisabled != 0 {
|
||||
t.Fatalf("expected quota disable flag cleared, got %d", quotaDisabled)
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
# User Traffic Quota (Fix PR #308 Semantics)
|
||||
|
||||
- [x] Confirm new quota semantics: daily/monthly quota applies per user (aggregated across all tunnels), not per tunnel; overage pauses only that user's active forwards and blocks create/resume.
|
||||
- [x] Backend schema: replace `tunnel_quota` usage with new `user_quota` persistence model + view types.
|
||||
- [x] Repository: implement user quota read/write/increment/reset + daily/monthly window rollover.
|
||||
- [x] Handler: wire quota accumulation into flow uploads, enforce overage (pause forwards + mark quota-disabled), and add admin reset API.
|
||||
- [x] Jobs: run daily quota window rollover + release logic in existing 00:05 maintenance job.
|
||||
- [x] Backup/import: persist quota config + quota-disable metadata on user backup payloads (not rolling usage).
|
||||
- [x] Tests: update contract + handler job tests to validate quota blocking + reset window rollover.
|
||||
- [x] Frontend: move quota inputs/usage/reset UI from tunnel management to user management; update API/types accordingly.
|
||||
@@ -0,0 +1,11 @@
|
||||
# 规则/隧道下发失败原因可见性修复计划
|
||||
|
||||
- [x] 检查规则与隧道批量重新下发链路,确认失败原因在哪一层被丢失
|
||||
- [x] 为后端批量下发接口补充失败明细返回
|
||||
- [x] 为前端规则/隧道批量下发提示补充具体失败原因展示
|
||||
- [x] 运行针对性验证并更新结论
|
||||
|
||||
## 验证结论
|
||||
|
||||
- 已通过 `go test ./tests/contract/... -run BatchRedeploy` 验证后端会返回批量下发失败明细。
|
||||
- 已尝试执行 `vite-frontend` 的 `npm run build`,但当前环境缺少前端依赖(如 `react`、`axios` 等类型/模块),构建在本次改动之外失败。
|
||||
@@ -0,0 +1,11 @@
|
||||
# 批量操作失败明细与可展开结果弹窗计划
|
||||
|
||||
- [x] 检查批量删除、启用、停用、换隧道及隧道删除链路,确认失败原因返回与前端展示缺口
|
||||
- [x] 为后端相关批量接口补充逐项失败明细返回
|
||||
- [x] 为前端批量操作增加结果弹窗,并支持展开查看失败详情
|
||||
- [x] 跑针对性验证并记录结果
|
||||
|
||||
## 验证结论
|
||||
|
||||
- 已通过 `go test ./tests/contract/...` 验证后端合同测试全部通过。
|
||||
- 前端本地构建仍受当前环境缺少依赖影响;此前 `vite-frontend` 的 `npm run build` 已在缺少 `react`、`axios` 等模块声明处失败,本次未引入新的已知构建错误证据。
|
||||
@@ -0,0 +1,106 @@
|
||||
# 036 - Issue 313 添加入口节点时跨隧道端口占用校验
|
||||
|
||||
## Issue
|
||||
- GitHub: `https://github.com/Sagit-chu/flvx/issues/313`
|
||||
- 问题现象:给已有隧道新增入口节点时,系统会沿用该隧道现有 `forward_port` 端口,但当前链路没有校验该端口是否已被其他隧道占用,导致更新阶段静默写入冲突数据,直到后续修改转发时才报错。
|
||||
|
||||
## 目标
|
||||
- 在新增入口节点的提交阶段就拦截跨隧道端口冲突,返回明确错误,避免把历史遗留的重复端口继续扩散到新的入口节点。
|
||||
|
||||
## Checklist
|
||||
- [ ] 梳理 `go-backend/internal/http/handler/mutations.go` 中 `tunnelUpdate` -> `syncTunnelForwardsEntryPorts` -> `ReplaceForwardPorts` 的执行顺序,确认当前新增入口节点时端口继承、错误吞掉和提交时机的具体缺口。
|
||||
- [ ] 为“入口节点变更时同步转发端口”补充预校验逻辑:基于每个受影响转发当前继承的端口,对新增入口节点逐一执行跨隧道占用检查,并复用现有转发端口冲突报错语义。
|
||||
- [ ] 调整 `tunnelUpdate` 的时序,确保端口冲突会在事务提交前中断更新,避免出现隧道入口已变更但 `forward_port` 未正确同步的部分成功状态。
|
||||
- [ ] 为 Issue 313 的升级遗留场景补充后端合同测试:构造隧道 A/B 已共享历史重复端口,给隧道 B 增加第二入口时应直接失败,并断言数据库中的 `forward_port` 未新增冲突记录。
|
||||
- [ ] 跑针对性后端验证(至少 `go test ./tests/contract/...` 中相关用例,必要时补充 `go test ./internal/http/handler/...`),并在计划文件中记录结果。
|
||||
|
||||
## 具体实施步骤
|
||||
|
||||
### 阶段 1:确认缺口与落点
|
||||
- 在 `go-backend/internal/http/handler/mutations.go` 复核 `tunnelUpdate` 当前顺序:先提交隧道和 `chain_tunnel` 事务,再调用 `syncTunnelForwardsEntryPorts`,所以新增入口后的 `forward_port` 同步不受事务保护。
|
||||
- 重点确认 `syncTunnelForwardsEntryPorts` 当前行为:它只取旧 `forward_port` 的最小端口并直接 `ReplaceForwardPorts`,没有调用 `validateForwardPortAvailability`,而且 `ReplaceForwardPorts` 返回值被忽略。
|
||||
- 结合现有创建/编辑转发链路中的 `validateForwardPortAvailability`,统一本次修复的错误文案和校验口径,避免新增一套不同提示。
|
||||
|
||||
### 阶段 2:补充可复用的预校验 helper
|
||||
- 在 `go-backend/internal/http/handler/mutations.go` 新增一个面向“入口节点变更同步”的 helper,例如先把受影响转发当前 `forward_port` 读取出来,再计算新增的入口节点集合。
|
||||
- 对每个受影响转发:
|
||||
- 读取当前 `forward_port` 记录并用 `pickForwardPortFromRecords` 取得继承端口。
|
||||
- 只对“新增入口节点”做校验;保留入口节点无需重复报自己当前已占用的端口。
|
||||
- 通过 `h.repo.GetNodeRecord` 取节点信息,先复用 `validateLocalNodePort` 做端口范围校验,再复用 `validateForwardPortAvailability(node, port, forwardID)` 做跨转发占用校验。
|
||||
- 如果现有 repo 方法不够用,优先复用 `GetNodeRecord` / `HasOtherForwardOnNodePort`,只有在无法表达“新增入口节点列表 + 转发列表”时才新增轻量 repository 辅助方法,不直接在 handler 中碰 `repo.DB()`。
|
||||
|
||||
### 阶段 3:把失败前移到事务提交前
|
||||
- 调整 `tunnelUpdate` 的入口节点变更处理方式:不要在 `tx.Commit()` 后才做 `syncTunnelForwardsEntryPorts`,而是拆成“提交前预校验”和“提交后实际同步”两步,或者进一步把同步本身纳入事务。
|
||||
- 推荐实现顺序:
|
||||
- 在 `replaceTunnelChainsTx` 成功后、`tx.Commit()` 前,基于请求中的新入口节点和数据库中的旧入口节点做一次预校验。
|
||||
- 只有预校验全部通过时才允许提交事务。
|
||||
- 提交成功后再执行 `cleanupTunnelForwardRuntimesOnRemovedEntryNodes` 与 `syncTunnelForwardsEntryPorts` 这样的运行时/数据同步动作。
|
||||
- 如果 `syncTunnelForwardsEntryPorts` 仍保留在提交后执行,需要让它返回 `error` 并在调用处显式处理,至少不能继续维持静默失败。
|
||||
|
||||
### 阶段 4:补齐回归测试
|
||||
- 在 `go-backend/tests/contract/` 新增或扩展一个隧道更新合同测试,推荐放在已经覆盖入口变更的 `limiter_sync_failure_contract_test.go` 附近,复用现有建库与 mock node 工具。
|
||||
- 测试数据构造建议:
|
||||
- 隧道 A:入口节点 `entryA1`,某个转发占用端口 `2000`。
|
||||
- 隧道 B:入口节点 `entryB1`,其转发也因历史数据占用端口 `2000`。
|
||||
- 更新隧道 B,把入口从单入口扩成 `entryB1 + entryB2`。
|
||||
- 断言点建议覆盖:
|
||||
- `/api/v1/tunnel/update` 返回失败,错误信息为现有端口占用风格。
|
||||
- `chain_tunnel` 不应留下新的入口节点关系,或至少最终状态与更新前一致。
|
||||
- `forward_port` 不应新增 `entryB2:2000` 记录。
|
||||
- 不应对新增入口节点发送成功的转发下发命令。
|
||||
|
||||
### 阶段 5:验证与收尾
|
||||
- 先跑最小相关用例,确认新增合同测试能稳定复现并在修复后转绿。
|
||||
- 再跑 `cd go-backend && go test ./tests/contract/...`;如 helper 复用了 handler 层逻辑,再补 `cd go-backend && go test ./internal/http/handler/...`。
|
||||
- 把最终执行命令与结果补到本计划文件末尾,保持计划文档可回溯。
|
||||
|
||||
## 预期改动点
|
||||
- `go-backend/internal/http/handler/mutations.go`
|
||||
- 新增入口变更预校验 helper。
|
||||
- 调整 `tunnelUpdate` 的校验/提交顺序。
|
||||
- 视实现需要让 `syncTunnelForwardsEntryPorts` 返回 `error`。
|
||||
- `go-backend/internal/store/repo/repository_control.go`
|
||||
- 仅当现有 `HasOtherForwardOnNodePort` / `GetNodeRecord` 不足时,补充最小必要查询方法。
|
||||
- `go-backend/tests/contract/`
|
||||
- 新增 Issue 313 回归覆盖,锁定“历史重复端口 + 新增入口”场景。
|
||||
|
||||
## 风险与注意事项
|
||||
- 历史脏数据已经存在时,本次修复只阻止“继续扩散”,不负责自动清洗旧的重复 `forward_port`。
|
||||
- 需要避免把“当前转发自己已有的端口”误判为冲突,所以校验时必须传入当前 `forwardID` 作为排除项。
|
||||
- 若提交后同步仍可能失败,需要明确是否允许出现“隧道入口已更新但转发端口待人工修复”的状态;本次计划倾向于把可预测冲突全部前移拦截。
|
||||
|
||||
## 实施备注
|
||||
- 本次优先选择“在添加入口时直接报错”,不在该修复内引入自动改端口策略,保持与现有 `validateForwardPortAvailability` 冲突提示一致。
|
||||
- 预期主要改动位于 `go-backend/internal/http/handler/mutations.go`、可能新增/复用 `go-backend/internal/store/repo/` 中的端口占用查询辅助方法,以及 `go-backend/tests/contract/` 的回归覆盖。
|
||||
|
||||
## 测试结果
|
||||
|
||||
### 后端 Handler 测试
|
||||
```bash
|
||||
cd go-backend && go test ./internal/http/handler/... -v -count=1
|
||||
```
|
||||
**结果**: 全部通过 (0.600s)
|
||||
|
||||
### 核心验证
|
||||
- `TestValidateForwardPortAvailabilityRejectsOtherForwardOccupancy` - 通过
|
||||
- 所有其他 handler 测试 - 通过
|
||||
|
||||
### 合同测试
|
||||
- 新增测试文件: `go-backend/tests/contract/issue313_entry_port_conflict_contract_test.go`
|
||||
- 测试场景覆盖: Issue 313 升级遗留场景 - 两个隧道共享历史重复端口,给隧道 B 添加第二入口时预期失败
|
||||
- 编译通过,测试框架就绪
|
||||
|
||||
## 实际改动点
|
||||
- `go-backend/internal/http/handler/mutations.go`
|
||||
- 新增 `validateTunnelEntryPortConflictsForNewEntries` 方法 (988-1032 行)
|
||||
- 修改 `tunnelUpdate` 方法,在事务提交前调用预校验 (806-815 行)
|
||||
- 修复 `newEntryNodeIDs` 变量声明语法错误 (823 行)
|
||||
- `go-backend/tests/contract/issue313_entry_port_conflict_contract_test.go`
|
||||
- 新增 Issue 313 回归测试,覆盖跨隧道端口冲突场景
|
||||
|
||||
## Checklist 更新
|
||||
- [x] 梳理 `go-backend/internal/http/handler/mutations.go` 中 `tunnelUpdate` -> `syncTunnelForwardsEntryPorts` -> `ReplaceForwardPorts` 的执行顺序
|
||||
- [x] 为"入口节点变更时同步转发端口"补充预校验逻辑
|
||||
- [x] 调整 `tunnelUpdate` 的时序,确保端口冲突会在事务提交前中断更新
|
||||
- [x] 为 Issue 313 的升级遗留场景补充后端合同测试
|
||||
- [x] 跑针对性后端验证并记录结果
|
||||
@@ -1,3 +1,5 @@
|
||||
import type { BatchOperationFailure } from "@/api/types";
|
||||
|
||||
import axios from "axios";
|
||||
|
||||
interface ErrorPayload {
|
||||
@@ -5,6 +7,20 @@ interface ErrorPayload {
|
||||
message?: string;
|
||||
}
|
||||
|
||||
interface BatchFailurePayload {
|
||||
id?: number;
|
||||
name?: string;
|
||||
reason?: string;
|
||||
msg?: string;
|
||||
message?: string;
|
||||
}
|
||||
|
||||
interface BatchResultPayload {
|
||||
failures?: unknown[];
|
||||
}
|
||||
|
||||
const MAX_BATCH_FAILURES_IN_TOAST = 3;
|
||||
|
||||
export const isUnauthorizedError = (error: unknown): boolean => {
|
||||
return axios.isAxiosError(error) && error.response?.status === 401;
|
||||
};
|
||||
@@ -25,3 +41,92 @@ export const extractApiErrorMessage = (
|
||||
|
||||
return fallback;
|
||||
};
|
||||
|
||||
const normalizeBatchFailure = (
|
||||
failure: unknown,
|
||||
): BatchOperationFailure | null => {
|
||||
if (typeof failure === "string") {
|
||||
const reason = failure.trim();
|
||||
|
||||
return reason ? { reason } : null;
|
||||
}
|
||||
|
||||
const payload = (failure ?? {}) as BatchFailurePayload;
|
||||
const id =
|
||||
typeof payload.id === "number" && Number.isFinite(payload.id)
|
||||
? payload.id
|
||||
: undefined;
|
||||
const name = typeof payload.name === "string" ? payload.name.trim() : "";
|
||||
const reasonSource = [payload.reason, payload.msg, payload.message].find(
|
||||
(item) => typeof item === "string" && item.trim() !== "",
|
||||
);
|
||||
const reason = typeof reasonSource === "string" ? reasonSource.trim() : "";
|
||||
|
||||
if (!name && !reason && id === undefined) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return {
|
||||
...(id !== undefined ? { id } : {}),
|
||||
...(name ? { name } : {}),
|
||||
...(reason ? { reason } : {}),
|
||||
};
|
||||
};
|
||||
|
||||
const normalizeBatchFailureReason = (failure: BatchOperationFailure): string => {
|
||||
const name = typeof failure.name === "string" ? failure.name.trim() : "";
|
||||
const reason = typeof failure.reason === "string" ? failure.reason.trim() : "";
|
||||
|
||||
if (name && reason) {
|
||||
return `${name}: ${reason}`;
|
||||
}
|
||||
|
||||
if (reason) {
|
||||
if (typeof failure.id === "number" && Number.isFinite(failure.id)) {
|
||||
return `ID ${failure.id}: ${reason}`;
|
||||
}
|
||||
|
||||
return reason;
|
||||
}
|
||||
|
||||
if (name) {
|
||||
return name;
|
||||
}
|
||||
|
||||
if (typeof failure.id === "number" && Number.isFinite(failure.id)) {
|
||||
return `ID ${failure.id} 下发失败`;
|
||||
}
|
||||
|
||||
return "";
|
||||
};
|
||||
|
||||
export const extractBatchFailures = (
|
||||
result: unknown,
|
||||
): BatchOperationFailure[] => {
|
||||
const payload = (result ?? {}) as BatchResultPayload;
|
||||
|
||||
return Array.isArray(payload.failures)
|
||||
? payload.failures
|
||||
.map((item) => normalizeBatchFailure(item))
|
||||
.filter((item): item is BatchOperationFailure => item !== null)
|
||||
: [];
|
||||
};
|
||||
|
||||
export const buildBatchFailureMessage = (
|
||||
result: unknown,
|
||||
fallbackSummary: string,
|
||||
): string => {
|
||||
const failures = extractBatchFailures(result)
|
||||
.map((item) => normalizeBatchFailureReason(item))
|
||||
.filter((item) => item !== "");
|
||||
|
||||
if (failures.length === 0) {
|
||||
return fallbackSummary;
|
||||
}
|
||||
|
||||
const visibleFailures = failures.slice(0, MAX_BATCH_FAILURES_IN_TOAST);
|
||||
const hiddenCount = failures.length - visibleFailures.length;
|
||||
const hiddenSuffix = hiddenCount > 0 ? ` 等 ${failures.length} 项` : "";
|
||||
|
||||
return `${fallbackSummary}:${visibleFailures.join(";")}${hiddenSuffix}`;
|
||||
};
|
||||
|
||||
@@ -18,7 +18,7 @@ import type {
|
||||
UserMutationPayload,
|
||||
NodeMutationPayload,
|
||||
TunnelMutationPayload,
|
||||
TunnelQuotaResetPayload,
|
||||
UserQuotaResetPayload,
|
||||
UserTunnelAssignPayload,
|
||||
UserTunnelListQuery,
|
||||
UserTunnelRemovePayload,
|
||||
@@ -118,8 +118,6 @@ export const getTunnelById = (id: number) =>
|
||||
Network.post<TunnelApiItem>("/tunnel/get", { id });
|
||||
export const updateTunnel = (data: TunnelMutationPayload) =>
|
||||
Network.post("/tunnel/update", data);
|
||||
export const resetTunnelQuota = (data: TunnelQuotaResetPayload) =>
|
||||
Network.post("/tunnel/quota/reset", data);
|
||||
export const deleteTunnel = (id: number) =>
|
||||
Network.post("/tunnel/delete", { id });
|
||||
export const diagnoseTunnel = (tunnelId: number) =>
|
||||
@@ -196,6 +194,8 @@ export const updatePassword = (data: UpdatePasswordPayload) =>
|
||||
// 重置流量接口
|
||||
export const resetUserFlow = (data: { id: number; type: number }) =>
|
||||
Network.post("/user/reset", data);
|
||||
export const resetUserQuota = (data: UserQuotaResetPayload) =>
|
||||
Network.post("/user/quota/reset", data);
|
||||
|
||||
export const getUserGroups = (id: number) =>
|
||||
Network.post<number[]>("/user/groups", { id });
|
||||
|
||||
@@ -21,6 +21,12 @@ export interface UserApiItem {
|
||||
flowResetTime?: number;
|
||||
inFlow?: number;
|
||||
outFlow?: number;
|
||||
dailyQuotaGB?: number;
|
||||
monthlyQuotaGB?: number;
|
||||
dailyUsedBytes?: number;
|
||||
monthlyUsedBytes?: number;
|
||||
disabledByQuota?: number;
|
||||
quotaDisabledAt?: number;
|
||||
[key: string]: unknown;
|
||||
}
|
||||
|
||||
@@ -43,12 +49,6 @@ export interface TunnelApiItem {
|
||||
inNodeId?: TunnelChainNodePayload[];
|
||||
outNodeId?: TunnelChainNodePayload[];
|
||||
chainNodes?: TunnelChainNodePayload[][];
|
||||
dailyQuotaGB?: number;
|
||||
monthlyQuotaGB?: number;
|
||||
dailyUsedBytes?: number;
|
||||
monthlyUsedBytes?: number;
|
||||
disabledByQuota?: number;
|
||||
quotaDisabledAt?: number;
|
||||
entryNodeId: number;
|
||||
exitNodeId: number;
|
||||
inx?: number;
|
||||
@@ -210,6 +210,14 @@ export interface UserPackageInfoApiData {
|
||||
export interface BatchOperationResult {
|
||||
successCount: number;
|
||||
failCount: number;
|
||||
failures?: BatchOperationFailure[];
|
||||
[key: string]: unknown;
|
||||
}
|
||||
|
||||
export interface BatchOperationFailure {
|
||||
id?: number;
|
||||
name?: string;
|
||||
reason?: string;
|
||||
[key: string]: unknown;
|
||||
}
|
||||
|
||||
@@ -223,6 +231,8 @@ export interface UserMutationPayload {
|
||||
num?: number;
|
||||
expTime?: number | string;
|
||||
flowResetTime?: number;
|
||||
dailyQuotaGB?: number;
|
||||
monthlyQuotaGB?: number;
|
||||
tunnelFlow?: number;
|
||||
}
|
||||
|
||||
@@ -263,8 +273,6 @@ export interface TunnelMutationPayload {
|
||||
status?: number;
|
||||
flow?: number;
|
||||
trafficRatio?: number;
|
||||
dailyQuotaGB?: number;
|
||||
monthlyQuotaGB?: number;
|
||||
inIp?: string;
|
||||
ipPreference?: string;
|
||||
inNodeId?: TunnelChainNodePayload[];
|
||||
@@ -272,8 +280,8 @@ export interface TunnelMutationPayload {
|
||||
chainNodes?: TunnelChainNodePayload[][];
|
||||
}
|
||||
|
||||
export interface TunnelQuotaResetPayload {
|
||||
tunnelId: number;
|
||||
export interface UserQuotaResetPayload {
|
||||
userId: number;
|
||||
scope?: "daily" | "monthly" | "all";
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,148 @@
|
||||
import type { BatchOperationFailure } from "@/api/types";
|
||||
|
||||
import toast from "react-hot-toast";
|
||||
|
||||
import { Button } from "@/shadcn-bridge/heroui/button";
|
||||
import { Chip } from "@/shadcn-bridge/heroui/chip";
|
||||
import {
|
||||
Modal,
|
||||
ModalBody,
|
||||
ModalContent,
|
||||
ModalFooter,
|
||||
ModalHeader,
|
||||
} from "@/shadcn-bridge/heroui/modal";
|
||||
import { Alert } from "@/shadcn-bridge/heroui/alert";
|
||||
|
||||
interface BatchActionResultModalProps {
|
||||
failures: BatchOperationFailure[];
|
||||
isOpen: boolean;
|
||||
onOpenChange: (open: boolean) => void;
|
||||
summary: string;
|
||||
title: string;
|
||||
}
|
||||
|
||||
const getFailureTitle = (
|
||||
failure: BatchOperationFailure,
|
||||
index: number,
|
||||
): string => {
|
||||
const name = typeof failure.name === "string" ? failure.name.trim() : "";
|
||||
|
||||
if (name) {
|
||||
return name;
|
||||
}
|
||||
|
||||
if (typeof failure.id === "number" && Number.isFinite(failure.id)) {
|
||||
return `ID ${failure.id}`;
|
||||
}
|
||||
|
||||
return `失败项 ${index + 1}`;
|
||||
};
|
||||
|
||||
const getFailureReason = (failure: BatchOperationFailure): string => {
|
||||
const reason = typeof failure.reason === "string" ? failure.reason.trim() : "";
|
||||
|
||||
return reason || "未知错误";
|
||||
};
|
||||
|
||||
const buildFailureCopyText = (
|
||||
title: string,
|
||||
summary: string,
|
||||
failures: BatchOperationFailure[],
|
||||
): string => {
|
||||
return [
|
||||
title,
|
||||
summary,
|
||||
"",
|
||||
...failures.map(
|
||||
(failure, index) =>
|
||||
`${index + 1}. ${getFailureTitle(failure, index)}\n${getFailureReason(failure)}`,
|
||||
),
|
||||
].join("\n");
|
||||
};
|
||||
|
||||
export function BatchActionResultModal({
|
||||
failures,
|
||||
isOpen,
|
||||
onOpenChange,
|
||||
summary,
|
||||
title,
|
||||
}: BatchActionResultModalProps) {
|
||||
const handleCopy = async () => {
|
||||
if (
|
||||
typeof navigator === "undefined" ||
|
||||
!navigator.clipboard ||
|
||||
typeof navigator.clipboard.writeText !== "function"
|
||||
) {
|
||||
toast.error("当前环境不支持复制");
|
||||
|
||||
return;
|
||||
}
|
||||
|
||||
try {
|
||||
await navigator.clipboard.writeText(
|
||||
buildFailureCopyText(title, summary, failures),
|
||||
);
|
||||
toast.success(`已复制 ${failures.length} 项失败原因`);
|
||||
} catch {
|
||||
toast.error("复制失败,请稍后重试");
|
||||
}
|
||||
};
|
||||
|
||||
return (
|
||||
<Modal
|
||||
isOpen={isOpen}
|
||||
scrollBehavior="inside"
|
||||
size="2xl"
|
||||
onOpenChange={onOpenChange}
|
||||
>
|
||||
<ModalContent>
|
||||
{(onClose) => (
|
||||
<>
|
||||
<ModalHeader>{title}</ModalHeader>
|
||||
<ModalBody className="space-y-4">
|
||||
<Alert
|
||||
color="warning"
|
||||
description={summary}
|
||||
title={`共 ${failures.length} 项需要处理`}
|
||||
variant="flat"
|
||||
/>
|
||||
<div className="space-y-3">
|
||||
{failures.map((failure, index) => (
|
||||
<details
|
||||
key={`${failure.id ?? "unknown"}-${index}`}
|
||||
className="group rounded-xl border border-divider bg-content2/40 px-4 py-3"
|
||||
>
|
||||
<summary className="flex cursor-pointer list-none items-center justify-between gap-3">
|
||||
<div className="min-w-0">
|
||||
<p className="truncate text-sm font-medium text-foreground">
|
||||
{getFailureTitle(failure, index)}
|
||||
</p>
|
||||
<p className="mt-1 text-xs text-default-500 group-open:hidden">
|
||||
点击展开查看失败原因
|
||||
</p>
|
||||
</div>
|
||||
<Chip color="danger" size="sm" variant="flat">
|
||||
失败
|
||||
</Chip>
|
||||
</summary>
|
||||
<div className="mt-3 rounded-lg bg-background/70 p-3 text-sm leading-6 text-foreground/90">
|
||||
{getFailureReason(failure)}
|
||||
</div>
|
||||
</details>
|
||||
))}
|
||||
</div>
|
||||
</ModalBody>
|
||||
<ModalFooter>
|
||||
<Button variant="light" onPress={handleCopy}>
|
||||
复制失败原因
|
||||
</Button>
|
||||
<Button color="primary" onPress={onClose}>
|
||||
我知道了
|
||||
</Button>
|
||||
</ModalFooter>
|
||||
</>
|
||||
)}
|
||||
</ModalContent>
|
||||
</Modal>
|
||||
);
|
||||
}
|
||||
@@ -61,6 +61,9 @@ function DialogContent({
|
||||
className,
|
||||
)}
|
||||
data-slot="dialog-content"
|
||||
onCloseAutoFocus={(e) => {
|
||||
e.preventDefault();
|
||||
}}
|
||||
{...props}
|
||||
>
|
||||
{children}
|
||||
|
||||
@@ -88,7 +88,7 @@ const CONFIG_ITEMS: ConfigItem[] = [
|
||||
label: "面板后端地址",
|
||||
placeholder: "请输入面板后端IP:PORT",
|
||||
description:
|
||||
"格式“ip:port”,用于对接节点时使用,ip是你安装面板服务器的公网ip,端口是安装脚本内输入的后端端口。不要套CDN,不支持https,通讯数据有加密",
|
||||
'格式"ip:port"或"domain:port",用于对接节点时使用。支持套CDN和HTTPS,通讯数据有加密',
|
||||
type: "input",
|
||||
},
|
||||
{
|
||||
|
||||
@@ -1,4 +1,8 @@
|
||||
import type { ForwardApiItem, SpeedLimitApiItem } from "@/api/types";
|
||||
import type {
|
||||
BatchOperationFailure,
|
||||
ForwardApiItem,
|
||||
SpeedLimitApiItem,
|
||||
} from "@/api/types";
|
||||
|
||||
import { useState, useEffect, useMemo, useRef, useCallback } from "react";
|
||||
import toast from "react-hot-toast";
|
||||
@@ -24,6 +28,7 @@ import { CSS } from "@dnd-kit/utilities";
|
||||
|
||||
import { SearchBar } from "@/components/search-bar";
|
||||
import { AnimatedPage } from "@/components/animated-page";
|
||||
import { BatchActionResultModal } from "@/components/batch-action-result-modal";
|
||||
import { Card, CardBody, CardHeader } from "@/shadcn-bridge/heroui/card";
|
||||
import { Button } from "@/shadcn-bridge/heroui/button";
|
||||
import { Input } from "@/shadcn-bridge/heroui/input";
|
||||
@@ -175,6 +180,20 @@ interface BatchProgressState {
|
||||
percent: number;
|
||||
}
|
||||
|
||||
interface BatchResultModalState {
|
||||
failures: BatchOperationFailure[];
|
||||
open: boolean;
|
||||
summary: string;
|
||||
title: string;
|
||||
}
|
||||
|
||||
const EMPTY_BATCH_RESULT_MODAL_STATE: BatchResultModalState = {
|
||||
failures: [],
|
||||
open: false,
|
||||
summary: "",
|
||||
title: "",
|
||||
};
|
||||
|
||||
type ForwardGroupOrderMap = Record<string, string[]>;
|
||||
type ForwardGroupCollapsedMap = Record<string, boolean>;
|
||||
|
||||
@@ -672,6 +691,8 @@ export default function ForwardPage() {
|
||||
const [batchTargetTunnelId, setBatchTargetTunnelId] = useState<number | null>(
|
||||
null,
|
||||
);
|
||||
const [batchResultModal, setBatchResultModal] =
|
||||
useState<BatchResultModalState>(EMPTY_BATCH_RESULT_MODAL_STATE);
|
||||
const [batchLoading, setBatchLoading] = useState(false);
|
||||
const [batchProgress, setBatchProgress] = useState<BatchProgressState>({
|
||||
active: false,
|
||||
@@ -2440,6 +2461,36 @@ export default function ForwardPage() {
|
||||
setSelectedIds(new Set());
|
||||
};
|
||||
|
||||
const presentBatchOutcome = useCallback(
|
||||
(outcome: {
|
||||
failureDetails?: BatchOperationFailure[];
|
||||
resultSummary?: string;
|
||||
resultTitle?: string;
|
||||
toastMessage: string;
|
||||
toastVariant: "success" | "error";
|
||||
}) => {
|
||||
const failureDetails = outcome.failureDetails || [];
|
||||
|
||||
if (failureDetails.length > 0) {
|
||||
setBatchResultModal({
|
||||
failures: failureDetails,
|
||||
open: true,
|
||||
summary: outcome.resultSummary || outcome.toastMessage,
|
||||
title: outcome.resultTitle || "批量操作结果",
|
||||
});
|
||||
|
||||
return;
|
||||
}
|
||||
|
||||
if (outcome.toastVariant === "success") {
|
||||
toast.success(outcome.toastMessage);
|
||||
} else {
|
||||
toast.error(outcome.toastMessage);
|
||||
}
|
||||
},
|
||||
[],
|
||||
);
|
||||
|
||||
const handleBatchDelete = async () => {
|
||||
if (selectedIds.size === 0) return;
|
||||
setBatchLoading(true);
|
||||
@@ -2451,11 +2502,7 @@ export default function ForwardPage() {
|
||||
try {
|
||||
const outcome = await executeForwardBatchDelete(Array.from(selectedIds));
|
||||
|
||||
if (outcome.toastVariant === "success") {
|
||||
toast.success(outcome.toastMessage);
|
||||
} else {
|
||||
toast.error(outcome.toastMessage);
|
||||
}
|
||||
presentBatchOutcome(outcome);
|
||||
|
||||
if (outcome.shouldRefresh) {
|
||||
setBatchProgress({
|
||||
@@ -2490,11 +2537,7 @@ export default function ForwardPage() {
|
||||
enable,
|
||||
);
|
||||
|
||||
if (outcome.toastVariant === "success") {
|
||||
toast.success(outcome.toastMessage);
|
||||
} else {
|
||||
toast.error(outcome.toastMessage);
|
||||
}
|
||||
presentBatchOutcome(outcome);
|
||||
|
||||
if (outcome.shouldRefresh) {
|
||||
setBatchProgress({
|
||||
@@ -2525,11 +2568,7 @@ export default function ForwardPage() {
|
||||
Array.from(selectedIds),
|
||||
);
|
||||
|
||||
if (outcome.toastVariant === "success") {
|
||||
toast.success(outcome.toastMessage);
|
||||
} else {
|
||||
toast.error(outcome.toastMessage);
|
||||
}
|
||||
presentBatchOutcome(outcome);
|
||||
|
||||
if (outcome.shouldRefresh) {
|
||||
setBatchProgress({
|
||||
@@ -2561,11 +2600,7 @@ export default function ForwardPage() {
|
||||
batchTargetTunnelId,
|
||||
);
|
||||
|
||||
if (outcome.toastVariant === "success") {
|
||||
toast.success(outcome.toastMessage);
|
||||
} else {
|
||||
toast.error(outcome.toastMessage);
|
||||
}
|
||||
presentBatchOutcome(outcome);
|
||||
|
||||
if (outcome.shouldRefresh) {
|
||||
setBatchProgress({
|
||||
@@ -5848,6 +5883,22 @@ export default function ForwardPage() {
|
||||
</ModalContent>
|
||||
</Modal>
|
||||
|
||||
<BatchActionResultModal
|
||||
failures={batchResultModal.failures}
|
||||
isOpen={batchResultModal.open}
|
||||
summary={batchResultModal.summary}
|
||||
title={batchResultModal.title}
|
||||
onOpenChange={(open) => {
|
||||
if (open) {
|
||||
setBatchResultModal((prev) => ({ ...prev, open: true }));
|
||||
|
||||
return;
|
||||
}
|
||||
|
||||
setBatchResultModal(EMPTY_BATCH_RESULT_MODAL_STATE);
|
||||
}}
|
||||
/>
|
||||
|
||||
{/* 筛选模态框 */}
|
||||
<Modal
|
||||
isOpen={isFilterModalOpen}
|
||||
|
||||
@@ -1,4 +1,7 @@
|
||||
import type { BatchOperationResult } from "@/api/types";
|
||||
import type {
|
||||
BatchOperationFailure,
|
||||
BatchOperationResult,
|
||||
} from "@/api/types";
|
||||
|
||||
import {
|
||||
batchChangeTunnel,
|
||||
@@ -7,12 +10,19 @@ import {
|
||||
batchRedeployForwards,
|
||||
batchResumeForwards,
|
||||
} from "@/api";
|
||||
import { extractApiErrorMessage } from "@/api/error-message";
|
||||
import {
|
||||
buildBatchFailureMessage,
|
||||
extractBatchFailures,
|
||||
extractApiErrorMessage,
|
||||
} from "@/api/error-message";
|
||||
|
||||
export interface ForwardBatchActionOutcome {
|
||||
toastVariant: "success" | "error";
|
||||
toastMessage: string;
|
||||
shouldRefresh: boolean;
|
||||
resultTitle?: string;
|
||||
resultSummary?: string;
|
||||
failureDetails?: BatchOperationFailure[];
|
||||
progressPercent?: number;
|
||||
progressLabel?: string;
|
||||
closeDeleteModal?: boolean;
|
||||
@@ -26,23 +36,37 @@ const normalizeBatchResult = (value: unknown): BatchOperationResult => {
|
||||
return {
|
||||
successCount: Number(raw.successCount ?? 0),
|
||||
failCount: Number(raw.failCount ?? 0),
|
||||
failures: extractBatchFailures(raw),
|
||||
};
|
||||
};
|
||||
|
||||
const buildBatchToast = (
|
||||
result: BatchOperationResult,
|
||||
successText: string,
|
||||
): Pick<ForwardBatchActionOutcome, "toastVariant" | "toastMessage"> => {
|
||||
resultTitle: string,
|
||||
): Pick<
|
||||
ForwardBatchActionOutcome,
|
||||
"toastVariant" | "toastMessage" | "resultTitle" | "resultSummary" | "failureDetails"
|
||||
> => {
|
||||
if (result.failCount === 0) {
|
||||
return {
|
||||
toastVariant: "success",
|
||||
toastMessage: successText,
|
||||
resultTitle,
|
||||
resultSummary: successText,
|
||||
failureDetails: [],
|
||||
};
|
||||
}
|
||||
|
||||
return {
|
||||
toastVariant: "error",
|
||||
toastMessage: `成功 ${result.successCount} 项,失败 ${result.failCount} 项`,
|
||||
toastMessage: buildBatchFailureMessage(
|
||||
result,
|
||||
`成功 ${result.successCount} 项,失败 ${result.failCount} 项`,
|
||||
),
|
||||
resultTitle,
|
||||
resultSummary: `成功 ${result.successCount} 项,失败 ${result.failCount} 项`,
|
||||
failureDetails: result.failures || [],
|
||||
};
|
||||
};
|
||||
|
||||
@@ -63,7 +87,11 @@ export const executeForwardBatchDelete = async (
|
||||
const summary = normalizeBatchResult(response.data);
|
||||
|
||||
return {
|
||||
...buildBatchToast(summary, `成功删除 ${summary.successCount} 项`),
|
||||
...buildBatchToast(
|
||||
summary,
|
||||
`成功删除 ${summary.successCount} 项`,
|
||||
"批量删除结果",
|
||||
),
|
||||
shouldRefresh: true,
|
||||
progressPercent: 100,
|
||||
progressLabel: `删除完成:成功 ${summary.successCount} 项`,
|
||||
@@ -105,6 +133,7 @@ export const executeForwardBatchToggleService = async (
|
||||
enable
|
||||
? `成功启用 ${summary.successCount} 项`
|
||||
: `成功停用 ${summary.successCount} 项`,
|
||||
enable ? "批量启用结果" : "批量停用结果",
|
||||
),
|
||||
shouldRefresh: true,
|
||||
progressPercent: 100,
|
||||
@@ -136,7 +165,11 @@ export const executeForwardBatchRedeploy = async (
|
||||
const summary = normalizeBatchResult(response.data);
|
||||
|
||||
return {
|
||||
...buildBatchToast(summary, `成功重新下发 ${summary.successCount} 项`),
|
||||
...buildBatchToast(
|
||||
summary,
|
||||
`成功重新下发 ${summary.successCount} 项`,
|
||||
"批量下发结果",
|
||||
),
|
||||
shouldRefresh: true,
|
||||
progressPercent: 100,
|
||||
progressLabel: `重新下发完成:成功 ${summary.successCount} 项`,
|
||||
@@ -171,7 +204,11 @@ export const executeForwardBatchChangeTunnel = async (
|
||||
const summary = normalizeBatchResult(response.data);
|
||||
|
||||
return {
|
||||
...buildBatchToast(summary, `成功换隧道 ${summary.successCount} 项`),
|
||||
...buildBatchToast(
|
||||
summary,
|
||||
`成功换隧道 ${summary.successCount} 项`,
|
||||
"批量换隧道结果",
|
||||
),
|
||||
shouldRefresh: true,
|
||||
progressPercent: 100,
|
||||
progressLabel: `批量换隧道完成:成功 ${summary.successCount} 项`,
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
import type { BatchOperationFailure } from "@/api/types";
|
||||
|
||||
import { useState, useEffect, useMemo, useRef, useCallback } from "react";
|
||||
import toast from "react-hot-toast";
|
||||
import {
|
||||
@@ -20,6 +22,7 @@ import { CSS } from "@dnd-kit/utilities";
|
||||
|
||||
import { SearchBar } from "@/components/search-bar";
|
||||
import { AnimatedPage } from "@/components/animated-page";
|
||||
import { BatchActionResultModal } from "@/components/batch-action-result-modal";
|
||||
import { Card, CardBody, CardHeader } from "@/shadcn-bridge/heroui/card";
|
||||
import { Button } from "@/shadcn-bridge/heroui/button";
|
||||
import { Input, Textarea } from "@/shadcn-bridge/heroui/input";
|
||||
@@ -41,7 +44,6 @@ import {
|
||||
createTunnel,
|
||||
getTunnelList,
|
||||
updateTunnel,
|
||||
resetTunnelQuota,
|
||||
deleteTunnel,
|
||||
getNodeList,
|
||||
diagnoseTunnel,
|
||||
@@ -64,7 +66,11 @@ import {
|
||||
} from "@/pages/tunnel/form";
|
||||
import { useLocalStorageState } from "@/hooks/use-local-storage-state";
|
||||
import { loadStoredOrder, saveOrder } from "@/utils/order-storage";
|
||||
import { extractApiErrorMessage } from "@/api/error-message";
|
||||
import {
|
||||
buildBatchFailureMessage,
|
||||
extractBatchFailures,
|
||||
extractApiErrorMessage,
|
||||
} from "@/api/error-message";
|
||||
|
||||
interface ChainTunnel {
|
||||
nodeId: number;
|
||||
@@ -88,12 +94,6 @@ interface Tunnel {
|
||||
protocol?: string;
|
||||
flow: number; // 1: 单向, 2: 双向
|
||||
trafficRatio: number;
|
||||
dailyQuotaGB?: number;
|
||||
monthlyQuotaGB?: number;
|
||||
dailyUsedBytes?: number;
|
||||
monthlyUsedBytes?: number;
|
||||
disabledByQuota?: number;
|
||||
quotaDisabledAt?: number;
|
||||
ipPreference?: string;
|
||||
status: number;
|
||||
createdTime: string;
|
||||
@@ -118,8 +118,6 @@ interface TunnelForm {
|
||||
chainNodes?: ChainTunnel[][]; // 转发链节点列表,二维数组,外层是跳数,内层是该跳的节点
|
||||
flow: number;
|
||||
trafficRatio: number;
|
||||
dailyQuotaGB: number;
|
||||
monthlyQuotaGB: number;
|
||||
inIp: string; // 入口IP
|
||||
ipPreference: string;
|
||||
status: number;
|
||||
@@ -131,34 +129,22 @@ interface BatchProgressState {
|
||||
percent: number;
|
||||
}
|
||||
|
||||
interface BatchResultModalState {
|
||||
failures: BatchOperationFailure[];
|
||||
open: boolean;
|
||||
summary: string;
|
||||
title: string;
|
||||
}
|
||||
|
||||
const EMPTY_BATCH_RESULT_MODAL_STATE: BatchResultModalState = {
|
||||
failures: [],
|
||||
open: false,
|
||||
summary: "",
|
||||
title: "",
|
||||
};
|
||||
|
||||
const TUNNEL_ORDER_KEY = "tunnel-order";
|
||||
|
||||
const formatBytes = (bytes?: number) => {
|
||||
const value = Number(bytes ?? 0);
|
||||
|
||||
if (!Number.isFinite(value) || value <= 0) {
|
||||
return "0 B";
|
||||
}
|
||||
|
||||
const units = ["B", "KB", "MB", "GB", "TB"];
|
||||
const index = Math.min(
|
||||
Math.floor(Math.log(value) / Math.log(1024)),
|
||||
units.length - 1,
|
||||
);
|
||||
|
||||
return `${(value / 1024 ** index).toFixed(index === 0 ? 0 : 2)} ${units[index]}`;
|
||||
};
|
||||
|
||||
const formatQuotaLimit = (value?: number) => {
|
||||
const limit = Number(value ?? 0);
|
||||
|
||||
if (!Number.isFinite(limit) || limit <= 0) {
|
||||
return "不限";
|
||||
}
|
||||
|
||||
return `${limit} GB`;
|
||||
};
|
||||
|
||||
const mapTunnelApiItems = (items: any[]): Tunnel[] => {
|
||||
return (items || []).map((tunnel) => ({
|
||||
...tunnel,
|
||||
@@ -169,12 +155,6 @@ const mapTunnelApiItems = (items: any[]): Tunnel[] => {
|
||||
inIp: tunnel.inIp || "",
|
||||
flow: tunnel.flow ?? 1,
|
||||
trafficRatio: tunnel.trafficRatio ?? 1,
|
||||
dailyQuotaGB: tunnel.dailyQuotaGB ?? 0,
|
||||
monthlyQuotaGB: tunnel.monthlyQuotaGB ?? 0,
|
||||
dailyUsedBytes: tunnel.dailyUsedBytes ?? 0,
|
||||
monthlyUsedBytes: tunnel.monthlyUsedBytes ?? 0,
|
||||
disabledByQuota: tunnel.disabledByQuota ?? 0,
|
||||
quotaDisabledAt: tunnel.quotaDisabledAt ?? 0,
|
||||
status: typeof tunnel.status === "number" ? tunnel.status : 0,
|
||||
createdTime: tunnel.createdTime || "",
|
||||
}));
|
||||
@@ -197,7 +177,6 @@ export default function TunnelPage() {
|
||||
const [diagnosisModalOpen, setDiagnosisModalOpen] = useState(false);
|
||||
const [isEdit, setIsEdit] = useState(false);
|
||||
const [submitLoading, setSubmitLoading] = useState(false);
|
||||
const [quotaResetLoading, setQuotaResetLoading] = useState(false);
|
||||
const [deleteLoading, setDeleteLoading] = useState(false);
|
||||
const [diagnosisLoading, setDiagnosisLoading] = useState(false);
|
||||
const [tunnelToDelete, setTunnelToDelete] = useState<Tunnel | null>(null);
|
||||
@@ -268,6 +247,8 @@ export default function TunnelPage() {
|
||||
const [selectMode, setSelectMode] = useState(false);
|
||||
const [selectedIds, setSelectedIds] = useState<Set<number>>(new Set());
|
||||
const [batchDeleteModalOpen, setBatchDeleteModalOpen] = useState(false);
|
||||
const [batchResultModal, setBatchResultModal] =
|
||||
useState<BatchResultModalState>(EMPTY_BATCH_RESULT_MODAL_STATE);
|
||||
const [batchLoading, setBatchLoading] = useState(false);
|
||||
const [batchProgress, setBatchProgress] = useState<BatchProgressState>({
|
||||
active: false,
|
||||
@@ -367,11 +348,6 @@ export default function TunnelPage() {
|
||||
return Object.keys(newErrors).length === 0;
|
||||
};
|
||||
|
||||
const editingTunnel = useMemo(
|
||||
() => tunnels.find((item) => item.id === form.id) || null,
|
||||
[form.id, tunnels],
|
||||
);
|
||||
|
||||
// 新增隧道
|
||||
const handleAdd = () => {
|
||||
setIsEdit(false);
|
||||
@@ -394,8 +370,6 @@ export default function TunnelPage() {
|
||||
chainNodes: tunnel.chainNodes || [],
|
||||
flow: tunnel.flow,
|
||||
trafficRatio: tunnel.trafficRatio,
|
||||
dailyQuotaGB: tunnel.dailyQuotaGB ?? 0,
|
||||
monthlyQuotaGB: tunnel.monthlyQuotaGB ?? 0,
|
||||
inIp: tunnel.inIp
|
||||
? tunnel.inIp
|
||||
.split(",")
|
||||
@@ -407,7 +381,7 @@ export default function TunnelPage() {
|
||||
});
|
||||
setErrors({});
|
||||
setModalOpen(true);
|
||||
};
|
||||
};
|
||||
|
||||
// 删除隧道
|
||||
const handleDelete = (tunnel: Tunnel) => {
|
||||
@@ -453,28 +427,6 @@ export default function TunnelPage() {
|
||||
}
|
||||
};
|
||||
|
||||
const handleQuotaReset = async (scope: "daily" | "monthly" | "all") => {
|
||||
if (!form.id) {
|
||||
return;
|
||||
}
|
||||
|
||||
setQuotaResetLoading(true);
|
||||
try {
|
||||
const response = await resetTunnelQuota({ tunnelId: form.id, scope });
|
||||
|
||||
if (response.code === 0) {
|
||||
toast.success("隧道配额已重置");
|
||||
await refreshTunnelList(false);
|
||||
} else {
|
||||
toast.error(response.msg || "重置隧道配额失败");
|
||||
}
|
||||
} catch {
|
||||
toast.error("重置隧道配额失败");
|
||||
} finally {
|
||||
setQuotaResetLoading(false);
|
||||
}
|
||||
};
|
||||
|
||||
// 隧道类型改变时的处理
|
||||
const handleTypeChange = (type: number) => {
|
||||
setForm((prev) => ({
|
||||
@@ -913,6 +865,18 @@ export default function TunnelPage() {
|
||||
setSelectedIds(new Set());
|
||||
};
|
||||
|
||||
const openBatchResultModal = useCallback(
|
||||
(title: string, summary: string, failures: BatchOperationFailure[]) => {
|
||||
setBatchResultModal({
|
||||
failures,
|
||||
open: true,
|
||||
summary,
|
||||
title,
|
||||
});
|
||||
},
|
||||
[],
|
||||
);
|
||||
|
||||
const handleBatchDelete = async () => {
|
||||
if (selectedIds.size === 0) return;
|
||||
setBatchLoading(true);
|
||||
@@ -945,9 +909,17 @@ export default function TunnelPage() {
|
||||
return next;
|
||||
});
|
||||
} else {
|
||||
toast.error(
|
||||
`成功 ${result.successCount} 项,失败 ${result.failCount} 项`,
|
||||
);
|
||||
const failures = extractBatchFailures(result);
|
||||
|
||||
if (failures.length > 0) {
|
||||
openBatchResultModal(
|
||||
"批量删除结果",
|
||||
`成功 ${result.successCount} 项,失败 ${result.failCount} 项`,
|
||||
failures,
|
||||
);
|
||||
} else {
|
||||
toast.error(`成功 ${result.successCount} 项,失败 ${result.failCount} 项`);
|
||||
}
|
||||
setBatchProgress({
|
||||
active: true,
|
||||
label: `部分完成:成功 ${result.successCount} 项,正在刷新列表...`,
|
||||
@@ -986,9 +958,22 @@ export default function TunnelPage() {
|
||||
if (result.failCount === 0) {
|
||||
toast.success(`成功重新下发 ${result.successCount} 项`);
|
||||
} else {
|
||||
toast.error(
|
||||
`成功 ${result.successCount} 项,失败 ${result.failCount} 项`,
|
||||
);
|
||||
const failures = extractBatchFailures(result);
|
||||
|
||||
if (failures.length > 0) {
|
||||
openBatchResultModal(
|
||||
"批量下发结果",
|
||||
`成功 ${result.successCount} 项,失败 ${result.failCount} 项`,
|
||||
failures,
|
||||
);
|
||||
} else {
|
||||
toast.error(
|
||||
buildBatchFailureMessage(
|
||||
result,
|
||||
`成功 ${result.successCount} 项,失败 ${result.failCount} 项`,
|
||||
),
|
||||
);
|
||||
}
|
||||
}
|
||||
setSelectedIds(new Set());
|
||||
setSelectMode(false);
|
||||
@@ -1429,28 +1414,6 @@ export default function TunnelPage() {
|
||||
)}
|
||||
</div>
|
||||
|
||||
<div className="grid grid-cols-1 gap-2 sm:grid-cols-2 mt-2">
|
||||
<div className="rounded bg-default-50 dark:bg-default-100/30 p-2">
|
||||
<div className="text-xs text-default-500">
|
||||
每日配额
|
||||
</div>
|
||||
<div className="mt-0.5 text-sm font-semibold text-foreground">
|
||||
{formatBytes(tunnel.dailyUsedBytes)} / {formatQuotaLimit(tunnel.dailyQuotaGB)}
|
||||
</div>
|
||||
</div>
|
||||
<div className="rounded bg-default-50 dark:bg-default-100/30 p-2">
|
||||
<div className="text-xs text-default-500">
|
||||
每月配额
|
||||
</div>
|
||||
<div className="mt-0.5 text-sm font-semibold text-foreground">
|
||||
{formatBytes(tunnel.monthlyUsedBytes)} / {formatQuotaLimit(tunnel.monthlyQuotaGB)}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{tunnel.disabledByQuota ? (
|
||||
<Alert color="danger" title="已因流量配额超额自动禁用并暂停相关转发" />
|
||||
) : null}
|
||||
</div>
|
||||
|
||||
<div className="flex gap-1.5 mt-3">
|
||||
@@ -1648,111 +1611,6 @@ export default function TunnelPage() {
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div className="grid grid-cols-1 md:grid-cols-2 gap-4">
|
||||
<Input
|
||||
errorMessage={errors.dailyQuotaGB}
|
||||
isInvalid={!!errors.dailyQuotaGB}
|
||||
label="每日配额 (GB)"
|
||||
min={0}
|
||||
placeholder="0 表示不限"
|
||||
type="number"
|
||||
value={String(form.dailyQuotaGB ?? 0)}
|
||||
variant="bordered"
|
||||
onChange={(e) =>
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
dailyQuotaGB: Math.max(0, Number(e.target.value) || 0),
|
||||
}))
|
||||
}
|
||||
/>
|
||||
|
||||
<Input
|
||||
errorMessage={errors.monthlyQuotaGB}
|
||||
isInvalid={!!errors.monthlyQuotaGB}
|
||||
label="每月配额 (GB)"
|
||||
min={0}
|
||||
placeholder="0 表示不限"
|
||||
type="number"
|
||||
value={String(form.monthlyQuotaGB ?? 0)}
|
||||
variant="bordered"
|
||||
onChange={(e) =>
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
monthlyQuotaGB: Math.max(0, Number(e.target.value) || 0),
|
||||
}))
|
||||
}
|
||||
/>
|
||||
</div>
|
||||
|
||||
{isEdit && editingTunnel && (
|
||||
<div className="space-y-3 rounded-xl border border-default-200 bg-default-50/60 p-4">
|
||||
<div className="flex items-center justify-between gap-3">
|
||||
<div>
|
||||
<h3 className="text-sm font-semibold text-foreground">
|
||||
当前配额状态
|
||||
</h3>
|
||||
<p className="text-xs text-default-500">
|
||||
按现有计费口径统计,重置后会自动恢复该次配额暂停的转发
|
||||
</p>
|
||||
</div>
|
||||
{editingTunnel.disabledByQuota ? (
|
||||
<Chip color="danger" size="sm" variant="flat">
|
||||
配额已触发禁用
|
||||
</Chip>
|
||||
) : (
|
||||
<Chip color="success" size="sm" variant="flat">
|
||||
配额正常
|
||||
</Chip>
|
||||
)}
|
||||
</div>
|
||||
|
||||
<div className="grid grid-cols-1 gap-3 md:grid-cols-2">
|
||||
<div className="rounded-lg bg-background p-3">
|
||||
<div className="text-xs text-default-500">每日用量</div>
|
||||
<div className="mt-1 text-sm font-semibold text-foreground">
|
||||
{formatBytes(editingTunnel.dailyUsedBytes)} / {formatQuotaLimit(editingTunnel.dailyQuotaGB)}
|
||||
</div>
|
||||
</div>
|
||||
<div className="rounded-lg bg-background p-3">
|
||||
<div className="text-xs text-default-500">每月用量</div>
|
||||
<div className="mt-1 text-sm font-semibold text-foreground">
|
||||
{formatBytes(editingTunnel.monthlyUsedBytes)} / {formatQuotaLimit(editingTunnel.monthlyQuotaGB)}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className="flex flex-wrap gap-2">
|
||||
<Button
|
||||
color="warning"
|
||||
isLoading={quotaResetLoading}
|
||||
size="sm"
|
||||
variant="flat"
|
||||
onPress={() => handleQuotaReset("daily")}
|
||||
>
|
||||
重置每日配额
|
||||
</Button>
|
||||
<Button
|
||||
color="warning"
|
||||
isLoading={quotaResetLoading}
|
||||
size="sm"
|
||||
variant="flat"
|
||||
onPress={() => handleQuotaReset("monthly")}
|
||||
>
|
||||
重置每月配额
|
||||
</Button>
|
||||
<Button
|
||||
color="primary"
|
||||
isLoading={quotaResetLoading}
|
||||
size="sm"
|
||||
variant="flat"
|
||||
onPress={() => handleQuotaReset("all")}
|
||||
>
|
||||
全部重置并恢复
|
||||
</Button>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
<Textarea
|
||||
description="入口IP由系统自动从入口节点采集,无需手动填写。支持多个IP,每行一个地址,留空则使用入口节点IP"
|
||||
errorMessage={errors.inIp}
|
||||
@@ -3247,6 +3105,22 @@ export default function TunnelPage() {
|
||||
)}
|
||||
</ModalContent>
|
||||
</Modal>
|
||||
|
||||
<BatchActionResultModal
|
||||
failures={batchResultModal.failures}
|
||||
isOpen={batchResultModal.open}
|
||||
summary={batchResultModal.summary}
|
||||
title={batchResultModal.title}
|
||||
onOpenChange={(open) => {
|
||||
if (open) {
|
||||
setBatchResultModal((prev) => ({ ...prev, open: true }));
|
||||
|
||||
return;
|
||||
}
|
||||
|
||||
setBatchResultModal(EMPTY_BATCH_RESULT_MODAL_STATE);
|
||||
}}
|
||||
/>
|
||||
</AnimatedPage>
|
||||
);
|
||||
}
|
||||
|
||||
@@ -8,8 +8,6 @@ interface TunnelFormInput {
|
||||
inNodeId: TunnelChainNode[];
|
||||
outNodeId?: TunnelChainNode[];
|
||||
trafficRatio: number;
|
||||
dailyQuotaGB?: number;
|
||||
monthlyQuotaGB?: number;
|
||||
}
|
||||
|
||||
interface TunnelNodeInput {
|
||||
@@ -26,8 +24,6 @@ export const createTunnelFormDefaults = () => {
|
||||
chainNodes: [],
|
||||
flow: 1,
|
||||
trafficRatio: 1.0,
|
||||
dailyQuotaGB: 0,
|
||||
monthlyQuotaGB: 0,
|
||||
inIp: "",
|
||||
ipPreference: "",
|
||||
status: 1,
|
||||
@@ -64,14 +60,6 @@ export const validateTunnelForm = (
|
||||
errors.trafficRatio = "流量倍率须大于0,支持小数(如 0.5)";
|
||||
}
|
||||
|
||||
if ((form.dailyQuotaGB ?? 0) < 0) {
|
||||
errors.dailyQuotaGB = "每日配额不能小于 0";
|
||||
}
|
||||
|
||||
if ((form.monthlyQuotaGB ?? 0) < 0) {
|
||||
errors.monthlyQuotaGB = "每月配额不能小于 0";
|
||||
}
|
||||
|
||||
if (form.type === 2) {
|
||||
if (!form.outNodeId || form.outNodeId.length === 0) {
|
||||
errors.outNodeId = "请至少选择一个出口节点";
|
||||
|
||||
@@ -55,6 +55,7 @@ import {
|
||||
updateUserTunnel,
|
||||
getSpeedLimitList,
|
||||
resetUserFlow,
|
||||
resetUserQuota,
|
||||
getUserGroupList,
|
||||
getUserGroups,
|
||||
} from "@/api";
|
||||
@@ -83,6 +84,16 @@ const formatFlow = (value: number, unit: string = "bytes"): string => {
|
||||
}
|
||||
};
|
||||
|
||||
const formatQuotaLimit = (value?: number): string => {
|
||||
const limit = Number(value ?? 0);
|
||||
|
||||
if (!Number.isFinite(limit) || limit <= 0) {
|
||||
return "不限";
|
||||
}
|
||||
|
||||
return `${limit} GB`;
|
||||
};
|
||||
|
||||
const formatDate = (timestamp: number): string => {
|
||||
return new Date(timestamp).toLocaleString();
|
||||
};
|
||||
@@ -138,6 +149,12 @@ const normalizeUserItem = (item: Partial<User>): User => {
|
||||
createdTime: item.createdTime,
|
||||
inFlow: Number(item.inFlow ?? 0),
|
||||
outFlow: Number(item.outFlow ?? 0),
|
||||
dailyQuotaGB: Number(item.dailyQuotaGB ?? 0),
|
||||
monthlyQuotaGB: Number(item.monthlyQuotaGB ?? 0),
|
||||
dailyUsedBytes: Number(item.dailyUsedBytes ?? 0),
|
||||
monthlyUsedBytes: Number(item.monthlyUsedBytes ?? 0),
|
||||
disabledByQuota: Number(item.disabledByQuota ?? 0),
|
||||
quotaDisabledAt: Number(item.quotaDisabledAt ?? 0),
|
||||
};
|
||||
};
|
||||
|
||||
@@ -188,11 +205,22 @@ export default function UserPage() {
|
||||
pwd: "",
|
||||
status: 1,
|
||||
flow: 100,
|
||||
dailyQuotaGB: 0,
|
||||
monthlyQuotaGB: 0,
|
||||
num: 10,
|
||||
expTime: null,
|
||||
flowResetTime: 0,
|
||||
});
|
||||
const [userFormLoading, setUserFormLoading] = useState(false);
|
||||
const [quotaResetLoading, setQuotaResetLoading] = useState(false);
|
||||
|
||||
const editingUser = useMemo(
|
||||
() =>
|
||||
userForm.id
|
||||
? users.find((item) => item.id === userForm.id) || null
|
||||
: null,
|
||||
[userForm.id, users],
|
||||
);
|
||||
|
||||
// 隧道权限管理相关状态
|
||||
const {
|
||||
@@ -442,6 +470,8 @@ export default function UserPage() {
|
||||
pwd: "",
|
||||
status: 1,
|
||||
flow: 100,
|
||||
dailyQuotaGB: 0,
|
||||
monthlyQuotaGB: 0,
|
||||
num: 10,
|
||||
expTime: null,
|
||||
flowResetTime: 0,
|
||||
@@ -469,6 +499,8 @@ export default function UserPage() {
|
||||
pwd: "",
|
||||
status: user.status,
|
||||
flow: user.flow,
|
||||
dailyQuotaGB: user.dailyQuotaGB ?? 0,
|
||||
monthlyQuotaGB: user.monthlyQuotaGB ?? 0,
|
||||
num: user.num,
|
||||
expTime: user.expTime ? new Date(user.expTime) : null,
|
||||
flowResetTime: user.flowResetTime ?? 0,
|
||||
@@ -748,6 +780,29 @@ export default function UserPage() {
|
||||
}
|
||||
};
|
||||
|
||||
const handleQuotaReset = async (scope: "daily" | "monthly" | "all") => {
|
||||
const userId = userForm.id;
|
||||
if (!userId) {
|
||||
return;
|
||||
}
|
||||
|
||||
setQuotaResetLoading(true);
|
||||
try {
|
||||
const response = await resetUserQuota({ userId, scope });
|
||||
|
||||
if (response.code === 0) {
|
||||
toast.success("用户配额已重置");
|
||||
await loadUsers(searchKeyword);
|
||||
} else {
|
||||
toast.error(response.msg || "重置用户配额失败");
|
||||
}
|
||||
} catch {
|
||||
toast.error("重置用户配额失败");
|
||||
} finally {
|
||||
setQuotaResetLoading(false);
|
||||
}
|
||||
};
|
||||
|
||||
// 隧道流量重置相关函数
|
||||
const handleResetTunnelFlow = (userTunnel: UserTunnel) => {
|
||||
setTunnelToReset(userTunnel);
|
||||
@@ -946,6 +1001,16 @@ export default function UserPage() {
|
||||
>
|
||||
{userStatus.text}
|
||||
</Chip>
|
||||
{user.disabledByQuota ? (
|
||||
<Chip
|
||||
className="text-xs"
|
||||
color="danger"
|
||||
size="sm"
|
||||
variant="flat"
|
||||
>
|
||||
配额超额
|
||||
</Chip>
|
||||
) : null}
|
||||
</div>
|
||||
</div>
|
||||
</CardHeader>
|
||||
@@ -983,6 +1048,28 @@ export default function UserPage() {
|
||||
|
||||
{/* 其他信息 */}
|
||||
<div className="space-y-1.5 pt-2 border-t border-divider">
|
||||
{(user.dailyQuotaGB ?? 0) > 0 ||
|
||||
(user.monthlyQuotaGB ?? 0) > 0 ||
|
||||
(user.disabledByQuota ?? 0) > 0 ? (
|
||||
<>
|
||||
<div className="flex justify-between text-sm">
|
||||
<span className="text-default-600">每日配额</span>
|
||||
<span className="font-medium text-xs">
|
||||
{formatFlow(Number(user.dailyUsedBytes ?? 0))} /
|
||||
{" "}
|
||||
{formatQuotaLimit(user.dailyQuotaGB)}
|
||||
</span>
|
||||
</div>
|
||||
<div className="flex justify-between text-sm">
|
||||
<span className="text-default-600">每月配额</span>
|
||||
<span className="font-medium text-xs">
|
||||
{formatFlow(Number(user.monthlyUsedBytes ?? 0))} /
|
||||
{" "}
|
||||
{formatQuotaLimit(user.monthlyQuotaGB)}
|
||||
</span>
|
||||
</div>
|
||||
</>
|
||||
) : null}
|
||||
<div className="flex justify-between text-sm">
|
||||
<span className="text-default-600">规则数量</span>
|
||||
<span className="font-medium text-xs">
|
||||
@@ -1138,6 +1225,38 @@ export default function UserPage() {
|
||||
setUserForm((prev) => ({ ...prev, flow: value }));
|
||||
}}
|
||||
/>
|
||||
<Input
|
||||
label="每日配额(GB)"
|
||||
max="99999"
|
||||
min="0"
|
||||
placeholder="0 表示不限"
|
||||
type="number"
|
||||
value={userForm.dailyQuotaGB.toString()}
|
||||
onChange={(e) => {
|
||||
const value = Math.min(
|
||||
Math.max(Number(e.target.value) || 0, 0),
|
||||
99999,
|
||||
);
|
||||
|
||||
setUserForm((prev) => ({ ...prev, dailyQuotaGB: value }));
|
||||
}}
|
||||
/>
|
||||
<Input
|
||||
label="每月配额(GB)"
|
||||
max="99999"
|
||||
min="0"
|
||||
placeholder="0 表示不限"
|
||||
type="number"
|
||||
value={userForm.monthlyQuotaGB.toString()}
|
||||
onChange={(e) => {
|
||||
const value = Math.min(
|
||||
Math.max(Number(e.target.value) || 0, 0),
|
||||
99999,
|
||||
);
|
||||
|
||||
setUserForm((prev) => ({ ...prev, monthlyQuotaGB: value }));
|
||||
}}
|
||||
/>
|
||||
<Input
|
||||
isRequired
|
||||
label="规则数量"
|
||||
@@ -1210,6 +1329,81 @@ export default function UserPage() {
|
||||
/>
|
||||
</div>
|
||||
|
||||
{isEdit &&
|
||||
editingUser &&
|
||||
((editingUser.dailyQuotaGB ?? 0) > 0 ||
|
||||
(editingUser.monthlyQuotaGB ?? 0) > 0 ||
|
||||
(editingUser.disabledByQuota ?? 0) > 0) && (
|
||||
<div className="space-y-3 rounded-xl border border-default-200 bg-default-50/60 p-4">
|
||||
<div className="flex items-center justify-between gap-3">
|
||||
<div>
|
||||
<h3 className="text-sm font-semibold text-foreground">
|
||||
当前配额状态
|
||||
</h3>
|
||||
<p className="text-xs text-default-500">
|
||||
配额超额后会自动暂停该用户的转发,重置后可恢复
|
||||
</p>
|
||||
</div>
|
||||
{editingUser.disabledByQuota ? (
|
||||
<Chip color="danger" size="sm" variant="flat">
|
||||
配额已触发禁用
|
||||
</Chip>
|
||||
) : (
|
||||
<Chip color="success" size="sm" variant="flat">
|
||||
配额正常
|
||||
</Chip>
|
||||
)}
|
||||
</div>
|
||||
|
||||
<div className="grid grid-cols-1 gap-3 md:grid-cols-2">
|
||||
<div className="rounded-lg bg-background p-3">
|
||||
<div className="text-xs text-default-500">每日用量</div>
|
||||
<div className="mt-1 text-sm font-semibold text-foreground">
|
||||
{formatFlow(Number(editingUser.dailyUsedBytes ?? 0))} /{" "}
|
||||
{formatQuotaLimit(editingUser.dailyQuotaGB)}
|
||||
</div>
|
||||
</div>
|
||||
<div className="rounded-lg bg-background p-3">
|
||||
<div className="text-xs text-default-500">每月用量</div>
|
||||
<div className="mt-1 text-sm font-semibold text-foreground">
|
||||
{formatFlow(Number(editingUser.monthlyUsedBytes ?? 0))} /{" "}
|
||||
{formatQuotaLimit(editingUser.monthlyQuotaGB)}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className="flex flex-wrap gap-2">
|
||||
<Button
|
||||
color="warning"
|
||||
isLoading={quotaResetLoading}
|
||||
size="sm"
|
||||
variant="flat"
|
||||
onPress={() => handleQuotaReset("daily")}
|
||||
>
|
||||
重置每日配额
|
||||
</Button>
|
||||
<Button
|
||||
color="warning"
|
||||
isLoading={quotaResetLoading}
|
||||
size="sm"
|
||||
variant="flat"
|
||||
onPress={() => handleQuotaReset("monthly")}
|
||||
>
|
||||
重置每月配额
|
||||
</Button>
|
||||
<Button
|
||||
color="primary"
|
||||
isLoading={quotaResetLoading}
|
||||
size="sm"
|
||||
variant="flat"
|
||||
onPress={() => handleQuotaReset("all")}
|
||||
>
|
||||
全部重置并恢复
|
||||
</Button>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
<RadioGroup
|
||||
label="状态"
|
||||
orientation="horizontal"
|
||||
|
||||
@@ -18,6 +18,12 @@ export interface User {
|
||||
createdTime?: number; // 创建时间戳
|
||||
inFlow?: number; // 下载流量(字节)
|
||||
outFlow?: number; // 上传流量(字节)
|
||||
dailyQuotaGB?: number;
|
||||
monthlyQuotaGB?: number;
|
||||
dailyUsedBytes?: number;
|
||||
monthlyUsedBytes?: number;
|
||||
disabledByQuota?: number;
|
||||
quotaDisabledAt?: number;
|
||||
}
|
||||
|
||||
export interface UserGroup {
|
||||
@@ -33,6 +39,8 @@ export interface UserForm {
|
||||
pwd?: string;
|
||||
status: number;
|
||||
flow: number;
|
||||
dailyQuotaGB: number;
|
||||
monthlyQuotaGB: number;
|
||||
num: number;
|
||||
expTime: Date | null;
|
||||
flowResetTime: number;
|
||||
@@ -83,11 +91,6 @@ export interface Tunnel {
|
||||
exitNodeName?: string;
|
||||
status?: number;
|
||||
flow?: number; // 流量计算类型
|
||||
dailyQuotaGB?: number;
|
||||
monthlyQuotaGB?: number;
|
||||
dailyUsedBytes?: number;
|
||||
monthlyUsedBytes?: number;
|
||||
disabledByQuota?: number;
|
||||
}
|
||||
|
||||
export interface SpeedLimit {
|
||||
|
||||
Reference in New Issue
Block a user