mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-02 08:56:38 +08:00
feat: finalize Go backend migration and deployment cutover
This commit is contained in:
@@ -0,0 +1,67 @@
|
||||
package handler
|
||||
|
||||
import "time"
|
||||
|
||||
const captchaTokenTTL = 5 * time.Minute
|
||||
|
||||
func (h *Handler) storeCaptchaToken(token string) {
|
||||
if h == nil {
|
||||
return
|
||||
}
|
||||
token = normalizeCaptchaToken(token)
|
||||
if token == "" {
|
||||
return
|
||||
}
|
||||
|
||||
h.captchaMu.Lock()
|
||||
defer h.captchaMu.Unlock()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
h.pruneExpiredCaptchaTokensLocked(now)
|
||||
h.captchaTokens[token] = now + int64(captchaTokenTTL/time.Millisecond)
|
||||
}
|
||||
|
||||
func (h *Handler) consumeCaptchaToken(token string) bool {
|
||||
if h == nil {
|
||||
return false
|
||||
}
|
||||
token = normalizeCaptchaToken(token)
|
||||
if token == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
h.captchaMu.Lock()
|
||||
defer h.captchaMu.Unlock()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
h.pruneExpiredCaptchaTokensLocked(now)
|
||||
expiresAt, ok := h.captchaTokens[token]
|
||||
if !ok || expiresAt <= now {
|
||||
delete(h.captchaTokens, token)
|
||||
return false
|
||||
}
|
||||
delete(h.captchaTokens, token)
|
||||
return true
|
||||
}
|
||||
|
||||
func (h *Handler) pruneExpiredCaptchaTokensLocked(now int64) {
|
||||
for token, expiresAt := range h.captchaTokens {
|
||||
if expiresAt <= now {
|
||||
delete(h.captchaTokens, token)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeCaptchaToken(token string) string {
|
||||
return trimToken(token)
|
||||
}
|
||||
|
||||
func trimToken(token string) string {
|
||||
for len(token) > 0 && (token[0] == ' ' || token[0] == '\t' || token[0] == '\n' || token[0] == '\r') {
|
||||
token = token[1:]
|
||||
}
|
||||
for len(token) > 0 && (token[len(token)-1] == ' ' || token[len(token)-1] == '\t' || token[len(token)-1] == '\n' || token[len(token)-1] == '\r') {
|
||||
token = token[:len(token)-1]
|
||||
}
|
||||
return token
|
||||
}
|
||||
@@ -27,9 +27,11 @@ type forwardRecord struct {
|
||||
}
|
||||
|
||||
type tunnelRecord struct {
|
||||
ID int64
|
||||
Type int
|
||||
Status int
|
||||
ID int64
|
||||
Type int
|
||||
Status int
|
||||
Flow int64
|
||||
TrafficRatio float64
|
||||
}
|
||||
|
||||
type forwardPortRecord struct {
|
||||
@@ -111,15 +113,21 @@ func (h *Handler) getForwardRecord(forwardID int64) (*forwardRecord, error) {
|
||||
}
|
||||
|
||||
func (h *Handler) getTunnelRecord(tunnelID int64) (*tunnelRecord, error) {
|
||||
row := h.repo.DB().QueryRow(`SELECT id, type, status FROM tunnel WHERE id = ? LIMIT 1`, tunnelID)
|
||||
row := h.repo.DB().QueryRow(`SELECT id, type, status, flow, traffic_ratio FROM tunnel WHERE id = ? LIMIT 1`, tunnelID)
|
||||
var tr tunnelRecord
|
||||
err := row.Scan(&tr.ID, &tr.Type, &tr.Status)
|
||||
err := row.Scan(&tr.ID, &tr.Type, &tr.Status, &tr.Flow, &tr.TrafficRatio)
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, errors.New("隧道不存在")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
if tr.Flow <= 0 {
|
||||
tr.Flow = 1
|
||||
}
|
||||
if tr.TrafficRatio <= 0 {
|
||||
tr.TrafficRatio = 1
|
||||
}
|
||||
return &tr, nil
|
||||
}
|
||||
|
||||
@@ -287,7 +295,7 @@ func (h *Handler) controlForwardServices(forward *forwardRecord, commandType str
|
||||
}
|
||||
base := buildForwardServiceBase(forward.ID, forward.UserID, userTunnelID)
|
||||
payload := map[string]interface{}{
|
||||
"services": []string{base + "_tcp", base + "_udp"},
|
||||
"services": []string{base, base + "_tcp", base + "_udp"},
|
||||
}
|
||||
seen := map[int64]struct{}{}
|
||||
for _, fp := range ports {
|
||||
|
||||
@@ -0,0 +1,338 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const bytesPerGB int64 = 1024 * 1024 * 1024
|
||||
|
||||
type userTunnelPolicy struct {
|
||||
ID int64
|
||||
UserID int64
|
||||
TunnelID int64
|
||||
Flow int64
|
||||
InFlow int64
|
||||
OutFlow int64
|
||||
ExpTime int64
|
||||
Status int
|
||||
}
|
||||
|
||||
type gostConfigSnapshot struct {
|
||||
Services []namedConfigItem `json:"services"`
|
||||
Chains []namedConfigItem `json:"chains"`
|
||||
Limiters []namedConfigItem `json:"limiters"`
|
||||
}
|
||||
|
||||
type namedConfigItem struct {
|
||||
Name string `json:"name"`
|
||||
}
|
||||
|
||||
func (h *Handler) processFlowItem(item flowItem) {
|
||||
serviceName := strings.TrimSpace(item.N)
|
||||
if serviceName == "" || serviceName == "web_api" {
|
||||
return
|
||||
}
|
||||
|
||||
forwardID, userID, userTunnelID, ok := parseFlowServiceIDs(serviceName)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
inFlow, outFlow := h.scaleFlowByTunnel(forwardID, item.D, item.U)
|
||||
_ = h.repo.AddFlow(forwardID, userID, userTunnelID, inFlow, outFlow)
|
||||
|
||||
if userTunnelID > 0 {
|
||||
h.enforceFlowPolicies(userID, userTunnelID)
|
||||
}
|
||||
}
|
||||
|
||||
func parseFlowServiceIDs(serviceName string) (int64, int64, int64, bool) {
|
||||
parts := strings.Split(serviceName, "_")
|
||||
if len(parts) < 3 {
|
||||
return 0, 0, 0, false
|
||||
}
|
||||
|
||||
forwardID, err1 := strconv.ParseInt(parts[0], 10, 64)
|
||||
userID, err2 := strconv.ParseInt(parts[1], 10, 64)
|
||||
userTunnelID, err3 := strconv.ParseInt(parts[2], 10, 64)
|
||||
if err1 != nil || err2 != nil || err3 != nil || forwardID <= 0 || userID <= 0 {
|
||||
return 0, 0, 0, false
|
||||
}
|
||||
|
||||
return forwardID, userID, userTunnelID, true
|
||||
}
|
||||
|
||||
func (h *Handler) scaleFlowByTunnel(forwardID int64, inFlow int64, outFlow int64) (int64, int64) {
|
||||
forward, err := h.getForwardRecord(forwardID)
|
||||
if err != nil || forward == nil {
|
||||
return inFlow, outFlow
|
||||
}
|
||||
|
||||
tunnel, err := h.getTunnelRecord(forward.TunnelID)
|
||||
if err != nil || tunnel == nil {
|
||||
return inFlow, outFlow
|
||||
}
|
||||
|
||||
scaledIn := int64(float64(inFlow)*tunnel.TrafficRatio) * tunnel.Flow
|
||||
scaledOut := int64(float64(outFlow)*tunnel.TrafficRatio) * tunnel.Flow
|
||||
return scaledIn, scaledOut
|
||||
}
|
||||
|
||||
func (h *Handler) enforceFlowPolicies(userID int64, userTunnelID int64) {
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
if h.shouldPauseUser(userID, now) {
|
||||
h.pauseUserForwards(userID, now)
|
||||
}
|
||||
|
||||
policy, err := h.getUserTunnelPolicy(userTunnelID)
|
||||
if err != nil || policy == nil {
|
||||
return
|
||||
}
|
||||
|
||||
if shouldPauseUserTunnel(policy, now) {
|
||||
h.pauseUserTunnelForwards(policy.UserID, policy.TunnelID, now)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) shouldPauseUser(userID int64, now int64) bool {
|
||||
user, err := h.repo.GetUserByID(userID)
|
||||
if err != nil || user == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
flowLimit := user.Flow * bytesPerGB
|
||||
current := user.InFlow + user.OutFlow
|
||||
if flowLimit < current {
|
||||
return true
|
||||
}
|
||||
if user.ExpTime > 0 && user.ExpTime <= now {
|
||||
return true
|
||||
}
|
||||
return user.Status != 1
|
||||
}
|
||||
|
||||
func shouldPauseUserTunnel(policy *userTunnelPolicy, now int64) bool {
|
||||
if policy == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
flowLimit := policy.Flow * bytesPerGB
|
||||
current := policy.InFlow + policy.OutFlow
|
||||
if current >= flowLimit {
|
||||
return true
|
||||
}
|
||||
if policy.ExpTime > 0 && policy.ExpTime <= now {
|
||||
return true
|
||||
}
|
||||
return policy.Status != 1
|
||||
}
|
||||
|
||||
func (h *Handler) getUserTunnelPolicy(userTunnelID int64) (*userTunnelPolicy, error) {
|
||||
if userTunnelID <= 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
row := h.repo.DB().QueryRow(`
|
||||
SELECT id, user_id, tunnel_id, flow, in_flow, out_flow, exp_time, status
|
||||
FROM user_tunnel
|
||||
WHERE id = ?
|
||||
LIMIT 1
|
||||
`, userTunnelID)
|
||||
|
||||
var policy userTunnelPolicy
|
||||
if err := row.Scan(&policy.ID, &policy.UserID, &policy.TunnelID, &policy.Flow, &policy.InFlow, &policy.OutFlow, &policy.ExpTime, &policy.Status); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return &policy, nil
|
||||
}
|
||||
|
||||
func (h *Handler) pauseUserForwards(userID int64, now int64) {
|
||||
forwards, err := h.listActiveForwardsByUser(userID)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
h.pauseForwardRecords(forwards, now)
|
||||
}
|
||||
|
||||
func (h *Handler) pauseUserTunnelForwards(userID int64, tunnelID int64, now int64) {
|
||||
forwards, err := h.listActiveForwardsByUserTunnel(userID, tunnelID)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
h.pauseForwardRecords(forwards, now)
|
||||
}
|
||||
|
||||
func (h *Handler) pauseForwardRecords(forwards []forwardRecord, now int64) {
|
||||
for i := range forwards {
|
||||
forward := forwards[i]
|
||||
_ = h.controlForwardServices(&forward, "PauseService", false)
|
||||
_, _ = h.repo.DB().Exec(`UPDATE forward SET status = 0, updated_time = ? WHERE id = ?`, now, forward.ID)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) listActiveForwardsByUser(userID int64) ([]forwardRecord, error) {
|
||||
rows, err := h.repo.DB().Query(`
|
||||
SELECT id, user_id, user_name, name, tunnel_id, remote_addr, strategy, status
|
||||
FROM forward
|
||||
WHERE user_id = ? AND status = 1
|
||||
ORDER BY id ASC
|
||||
`, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
return scanForwardRecords(rows)
|
||||
}
|
||||
|
||||
func (h *Handler) listActiveForwardsByUserTunnel(userID int64, tunnelID int64) ([]forwardRecord, error) {
|
||||
rows, err := h.repo.DB().Query(`
|
||||
SELECT id, user_id, user_name, name, tunnel_id, remote_addr, strategy, status
|
||||
FROM forward
|
||||
WHERE user_id = ? AND tunnel_id = ? AND status = 1
|
||||
ORDER BY id ASC
|
||||
`, userID, tunnelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
return scanForwardRecords(rows)
|
||||
}
|
||||
|
||||
func scanForwardRecords(rows *sql.Rows) ([]forwardRecord, error) {
|
||||
out := make([]forwardRecord, 0)
|
||||
for rows.Next() {
|
||||
var record forwardRecord
|
||||
if err := rows.Scan(&record.ID, &record.UserID, &record.UserName, &record.Name, &record.TunnelID, &record.RemoteAddr, &record.Strategy, &record.Status); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if strings.TrimSpace(record.Strategy) == "" {
|
||||
record.Strategy = "fifo"
|
||||
}
|
||||
out = append(out, record)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (h *Handler) cleanNodeConfigs(nodeID int64, rawConfig string) {
|
||||
if h == nil || h.repo == nil || h.repo.DB() == nil || nodeID <= 0 {
|
||||
return
|
||||
}
|
||||
if strings.TrimSpace(rawConfig) == "" {
|
||||
return
|
||||
}
|
||||
|
||||
var snapshot gostConfigSnapshot
|
||||
if err := json.Unmarshal([]byte(rawConfig), &snapshot); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
h.cleanOrphanedServices(nodeID, snapshot.Services)
|
||||
h.cleanOrphanedChains(nodeID, snapshot.Chains)
|
||||
h.cleanOrphanedLimiters(nodeID, snapshot.Limiters)
|
||||
}
|
||||
|
||||
func (h *Handler) cleanOrphanedServices(nodeID int64, services []namedConfigItem) {
|
||||
for _, item := range services {
|
||||
name := strings.TrimSpace(item.Name)
|
||||
if name == "" || name == "web_api" {
|
||||
continue
|
||||
}
|
||||
|
||||
parts := strings.Split(name, "_")
|
||||
if len(parts) >= 3 {
|
||||
forwardID, err := strconv.ParseInt(parts[0], 10, 64)
|
||||
if err == nil && forwardID > 0 && !h.forwardExists(forwardID) {
|
||||
_, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": []string{name, parts[0] + "_" + parts[1] + "_" + parts[2], parts[0] + "_" + parts[1] + "_" + parts[2] + "_tcp", parts[0] + "_" + parts[1] + "_" + parts[2] + "_udp"}}, false, true)
|
||||
continue
|
||||
}
|
||||
}
|
||||
suffix := parts[len(parts)-1]
|
||||
|
||||
switch suffix {
|
||||
case "tls":
|
||||
tunnelID, err := strconv.ParseInt(parts[0], 10, 64)
|
||||
if err != nil || tunnelID <= 0 || h.tunnelExists(tunnelID) {
|
||||
continue
|
||||
}
|
||||
_, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": []string{name}}, false, true)
|
||||
case "tcp":
|
||||
if len(parts) < 4 {
|
||||
continue
|
||||
}
|
||||
forwardID, err := strconv.ParseInt(parts[0], 10, 64)
|
||||
if err != nil || forwardID <= 0 || h.forwardExists(forwardID) {
|
||||
continue
|
||||
}
|
||||
base := strings.TrimSuffix(name, "_tcp")
|
||||
_, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": []string{base + "_tcp", base + "_udp"}}, false, true)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) cleanOrphanedChains(nodeID int64, chains []namedConfigItem) {
|
||||
for _, item := range chains {
|
||||
name := strings.TrimSpace(item.Name)
|
||||
if name == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
idx := strings.LastIndex(name, "_")
|
||||
if idx <= 0 || idx >= len(name)-1 {
|
||||
continue
|
||||
}
|
||||
tunnelID, err := strconv.ParseInt(name[idx+1:], 10, 64)
|
||||
if err != nil || tunnelID <= 0 || h.tunnelExists(tunnelID) {
|
||||
continue
|
||||
}
|
||||
_, _ = h.sendNodeCommand(nodeID, "DeleteChains", map[string]interface{}{"chain": name}, false, true)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) cleanOrphanedLimiters(nodeID int64, limiters []namedConfigItem) {
|
||||
for _, item := range limiters {
|
||||
name := strings.TrimSpace(item.Name)
|
||||
if name == "" || h.speedLimiterExists(name) {
|
||||
continue
|
||||
}
|
||||
_, _ = h.sendNodeCommand(nodeID, "DeleteLimiters", map[string]interface{}{"limiter": name}, false, true)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) tunnelExists(tunnelID int64) bool {
|
||||
var count int
|
||||
err := h.repo.DB().QueryRow(`SELECT COUNT(1) FROM tunnel WHERE id = ?`, tunnelID).Scan(&count)
|
||||
return err == nil && count > 0
|
||||
}
|
||||
|
||||
func (h *Handler) forwardExists(forwardID int64) bool {
|
||||
var count int
|
||||
err := h.repo.DB().QueryRow(`SELECT COUNT(1) FROM forward WHERE id = ?`, forwardID).Scan(&count)
|
||||
return err == nil && count > 0
|
||||
}
|
||||
|
||||
func (h *Handler) speedLimiterExists(name string) bool {
|
||||
if name == "" {
|
||||
return false
|
||||
}
|
||||
id, err := strconv.ParseInt(name, 10, 64)
|
||||
if err != nil || id <= 0 {
|
||||
return false
|
||||
}
|
||||
|
||||
var count int
|
||||
err = h.repo.DB().QueryRow(`SELECT COUNT(1) FROM speed_limit WHERE id = ?`, id).Scan(&count)
|
||||
return err == nil && count > 0
|
||||
}
|
||||
@@ -1,6 +1,7 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
@@ -9,6 +10,7 @@ import (
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
@@ -23,6 +25,14 @@ type Handler struct {
|
||||
repo *sqlite.Repository
|
||||
jwtSecret string
|
||||
wsServer *ws.Server
|
||||
|
||||
captchaMu sync.Mutex
|
||||
captchaTokens map[string]int64
|
||||
|
||||
jobsMu sync.Mutex
|
||||
jobsCancel context.CancelFunc
|
||||
jobsStarted bool
|
||||
jobsWG sync.WaitGroup
|
||||
}
|
||||
|
||||
type loginRequest struct {
|
||||
@@ -54,7 +64,12 @@ type flowItem struct {
|
||||
}
|
||||
|
||||
func New(repo *sqlite.Repository, jwtSecret string) *Handler {
|
||||
return &Handler{repo: repo, jwtSecret: jwtSecret, wsServer: ws.NewServer(repo, jwtSecret)}
|
||||
return &Handler{
|
||||
repo: repo,
|
||||
jwtSecret: jwtSecret,
|
||||
wsServer: ws.NewServer(repo, jwtSecret),
|
||||
captchaTokens: make(map[string]int64),
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) WebSocketHandler() http.Handler {
|
||||
@@ -170,6 +185,10 @@ func (h *Handler) login(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("验证码校验失败"))
|
||||
return
|
||||
}
|
||||
if captchaEnabled && !h.consumeCaptchaToken(req.CaptchaID) {
|
||||
response.WriteJSON(w, response.ErrDefault("验证码校验失败"))
|
||||
return
|
||||
}
|
||||
|
||||
user, err := h.repo.GetUserByUsername(req.Username)
|
||||
if err != nil {
|
||||
@@ -558,13 +577,17 @@ func (h *Handler) flowTest(w http.ResponseWriter, _ *http.Request) {
|
||||
|
||||
func (h *Handler) flowConfig(w http.ResponseWriter, r *http.Request) {
|
||||
secret := r.URL.Query().Get("secret")
|
||||
if ok, _ := h.repo.NodeExistsBySecret(secret); !ok {
|
||||
node, err := h.repo.GetNodeBySecret(secret)
|
||||
if err != nil || node == nil {
|
||||
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
|
||||
_, _ = w.Write([]byte("ok"))
|
||||
return
|
||||
}
|
||||
|
||||
_, _ = readAndDecryptFlowBody(r.Body, secret)
|
||||
rawData, err := readAndDecryptFlowBody(r.Body, secret)
|
||||
if err == nil && strings.TrimSpace(rawData) != "" {
|
||||
h.cleanNodeConfigs(node.ID, rawData)
|
||||
}
|
||||
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
|
||||
_, _ = w.Write([]byte("ok"))
|
||||
}
|
||||
@@ -582,17 +605,7 @@ func (h *Handler) flowUpload(w http.ResponseWriter, r *http.Request) {
|
||||
var items []flowItem
|
||||
if json.Unmarshal([]byte(raw), &items) == nil {
|
||||
for _, item := range items {
|
||||
parts := strings.Split(item.N, "_")
|
||||
if len(parts) < 3 || item.N == "web_api" {
|
||||
continue
|
||||
}
|
||||
forwardID, err1 := strconv.ParseInt(parts[0], 10, 64)
|
||||
userID, err2 := strconv.ParseInt(parts[1], 10, 64)
|
||||
userTunnelID, err3 := strconv.ParseInt(parts[2], 10, 64)
|
||||
if err1 != nil || err2 != nil || err3 != nil {
|
||||
continue
|
||||
}
|
||||
_ = h.repo.AddFlow(forwardID, userID, userTunnelID, item.D, item.U)
|
||||
h.processFlowItem(item)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,270 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"time"
|
||||
)
|
||||
|
||||
func (h *Handler) StartBackgroundJobs() {
|
||||
if h == nil || h.repo == nil || h.repo.DB() == nil {
|
||||
return
|
||||
}
|
||||
|
||||
h.jobsMu.Lock()
|
||||
if h.jobsStarted {
|
||||
h.jobsMu.Unlock()
|
||||
return
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
h.jobsCancel = cancel
|
||||
h.jobsStarted = true
|
||||
h.jobsWG.Add(2)
|
||||
h.jobsMu.Unlock()
|
||||
|
||||
go h.runHourlyStatsLoop(ctx)
|
||||
go h.runDailyMaintenanceLoop(ctx)
|
||||
}
|
||||
|
||||
func (h *Handler) StopBackgroundJobs() {
|
||||
if h == nil {
|
||||
return
|
||||
}
|
||||
|
||||
h.jobsMu.Lock()
|
||||
if !h.jobsStarted {
|
||||
h.jobsMu.Unlock()
|
||||
return
|
||||
}
|
||||
cancel := h.jobsCancel
|
||||
h.jobsCancel = nil
|
||||
h.jobsStarted = false
|
||||
h.jobsMu.Unlock()
|
||||
|
||||
if cancel != nil {
|
||||
cancel()
|
||||
}
|
||||
h.jobsWG.Wait()
|
||||
}
|
||||
|
||||
func (h *Handler) runHourlyStatsLoop(ctx context.Context) {
|
||||
defer h.jobsWG.Done()
|
||||
|
||||
for {
|
||||
wait := durationUntilNextHour(time.Now())
|
||||
timer := time.NewTimer(wait)
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
if !timer.Stop() {
|
||||
<-timer.C
|
||||
}
|
||||
return
|
||||
case <-timer.C:
|
||||
h.runStatisticsFlowJob(time.Now())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) runDailyMaintenanceLoop(ctx context.Context) {
|
||||
defer h.jobsWG.Done()
|
||||
|
||||
for {
|
||||
wait := durationUntilNextDailyMaintenance(time.Now())
|
||||
timer := time.NewTimer(wait)
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
if !timer.Stop() {
|
||||
<-timer.C
|
||||
}
|
||||
return
|
||||
case <-timer.C:
|
||||
h.runResetAndExpiryJob(time.Now())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func durationUntilNextHour(now time.Time) time.Duration {
|
||||
next := now.Truncate(time.Hour).Add(time.Hour)
|
||||
return next.Sub(now)
|
||||
}
|
||||
|
||||
func durationUntilNextDailyMaintenance(now time.Time) time.Duration {
|
||||
next := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 5, 0, now.Location())
|
||||
if !next.After(now) {
|
||||
next = next.Add(24 * time.Hour)
|
||||
}
|
||||
return next.Sub(now)
|
||||
}
|
||||
|
||||
func (h *Handler) runStatisticsFlowJob(now time.Time) {
|
||||
if h == nil || h.repo == nil || h.repo.DB() == nil {
|
||||
return
|
||||
}
|
||||
|
||||
db := h.repo.DB()
|
||||
nowMs := now.UnixMilli()
|
||||
cutoffMs := nowMs - int64((48*time.Hour)/time.Millisecond)
|
||||
_, _ = db.Exec(`DELETE FROM statistics_flow WHERE created_time < ?`, cutoffMs)
|
||||
|
||||
hourMark := now.Truncate(time.Hour)
|
||||
hourText := hourMark.Format("15:04")
|
||||
createdTime := hourMark.UnixMilli()
|
||||
|
||||
rows, err := db.Query(`SELECT id, in_flow, out_flow FROM user ORDER BY id ASC`)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
type userFlowSnapshot struct {
|
||||
userID int64
|
||||
inFlow int64
|
||||
outFlow int64
|
||||
}
|
||||
users := make([]userFlowSnapshot, 0)
|
||||
|
||||
for rows.Next() {
|
||||
var userID int64
|
||||
var inFlow int64
|
||||
var outFlow int64
|
||||
if err := rows.Scan(&userID, &inFlow, &outFlow); err != nil {
|
||||
continue
|
||||
}
|
||||
users = append(users, userFlowSnapshot{userID: userID, inFlow: inFlow, outFlow: outFlow})
|
||||
}
|
||||
_ = rows.Close()
|
||||
|
||||
for _, user := range users {
|
||||
currentTotal := user.inFlow + user.outFlow
|
||||
increment := currentTotal
|
||||
|
||||
var lastTotal sql.NullInt64
|
||||
err := db.QueryRow(`SELECT total_flow FROM statistics_flow WHERE user_id = ? ORDER BY id DESC LIMIT 1`, user.userID).Scan(&lastTotal)
|
||||
if err == nil && lastTotal.Valid {
|
||||
increment = currentTotal - lastTotal.Int64
|
||||
if increment < 0 {
|
||||
increment = currentTotal
|
||||
}
|
||||
}
|
||||
|
||||
_, _ = db.Exec(`
|
||||
INSERT INTO statistics_flow(user_id, flow, total_flow, time, created_time)
|
||||
VALUES(?, ?, ?, ?, ?)
|
||||
`, user.userID, increment, currentTotal, hourText, createdTime)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) runResetAndExpiryJob(now time.Time) {
|
||||
if h == nil || h.repo == nil || h.repo.DB() == nil {
|
||||
return
|
||||
}
|
||||
|
||||
h.resetMonthlyFlow(now)
|
||||
h.disableExpiredUsers(now.UnixMilli())
|
||||
h.disableExpiredUserTunnels(now.UnixMilli())
|
||||
}
|
||||
|
||||
func (h *Handler) resetMonthlyFlow(now time.Time) {
|
||||
db := h.repo.DB()
|
||||
currentDay := now.Day()
|
||||
lastDay := time.Date(now.Year(), now.Month()+1, 0, 0, 0, 0, 0, now.Location()).Day()
|
||||
|
||||
if currentDay == lastDay {
|
||||
_, _ = db.Exec(`
|
||||
UPDATE user
|
||||
SET in_flow = 0, out_flow = 0
|
||||
WHERE flow_reset_time != 0
|
||||
AND (flow_reset_time = ? OR flow_reset_time > ?)
|
||||
`, currentDay, lastDay)
|
||||
_, _ = db.Exec(`
|
||||
UPDATE user_tunnel
|
||||
SET in_flow = 0, out_flow = 0
|
||||
WHERE flow_reset_time != 0
|
||||
AND (flow_reset_time = ? OR flow_reset_time > ?)
|
||||
`, currentDay, lastDay)
|
||||
return
|
||||
}
|
||||
|
||||
_, _ = db.Exec(`
|
||||
UPDATE user
|
||||
SET in_flow = 0, out_flow = 0
|
||||
WHERE flow_reset_time != 0
|
||||
AND flow_reset_time = ?
|
||||
`, currentDay)
|
||||
_, _ = db.Exec(`
|
||||
UPDATE user_tunnel
|
||||
SET in_flow = 0, out_flow = 0
|
||||
WHERE flow_reset_time != 0
|
||||
AND flow_reset_time = ?
|
||||
`, currentDay)
|
||||
}
|
||||
|
||||
func (h *Handler) disableExpiredUsers(nowMs int64) {
|
||||
db := h.repo.DB()
|
||||
rows, err := db.Query(`
|
||||
SELECT id
|
||||
FROM user
|
||||
WHERE role_id != 0
|
||||
AND status = 1
|
||||
AND exp_time IS NOT NULL
|
||||
AND exp_time < ?
|
||||
`, nowMs)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
userIDs := make([]int64, 0)
|
||||
|
||||
for rows.Next() {
|
||||
var userID int64
|
||||
if err := rows.Scan(&userID); err != nil {
|
||||
continue
|
||||
}
|
||||
userIDs = append(userIDs, userID)
|
||||
}
|
||||
_ = rows.Close()
|
||||
|
||||
for _, userID := range userIDs {
|
||||
forwards, err := h.listActiveForwardsByUser(userID)
|
||||
if err == nil {
|
||||
h.pauseForwardRecords(forwards, nowMs)
|
||||
}
|
||||
_, _ = db.Exec(`UPDATE user SET status = 0 WHERE id = ?`, userID)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) disableExpiredUserTunnels(nowMs int64) {
|
||||
db := h.repo.DB()
|
||||
rows, err := db.Query(`
|
||||
SELECT id, user_id, tunnel_id
|
||||
FROM user_tunnel
|
||||
WHERE status = 1
|
||||
AND exp_time IS NOT NULL
|
||||
AND exp_time < ?
|
||||
`, nowMs)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
type expiredUserTunnel struct {
|
||||
userTunnelID int64
|
||||
userID int64
|
||||
tunnelID int64
|
||||
}
|
||||
items := make([]expiredUserTunnel, 0)
|
||||
|
||||
for rows.Next() {
|
||||
var userTunnelID int64
|
||||
var userID int64
|
||||
var tunnelID int64
|
||||
if err := rows.Scan(&userTunnelID, &userID, &tunnelID); err != nil {
|
||||
continue
|
||||
}
|
||||
items = append(items, expiredUserTunnel{userTunnelID: userTunnelID, userID: userID, tunnelID: tunnelID})
|
||||
}
|
||||
_ = rows.Close()
|
||||
|
||||
for _, item := range items {
|
||||
forwards, err := h.listActiveForwardsByUserTunnel(item.userID, item.tunnelID)
|
||||
if err == nil {
|
||||
h.pauseForwardRecords(forwards, nowMs)
|
||||
}
|
||||
_, _ = db.Exec(`UPDATE user_tunnel SET status = 0 WHERE id = ?`, item.userTunnelID)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,128 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/sqlite"
|
||||
)
|
||||
|
||||
func TestRunStatisticsFlowJobTracksIncrementAndPrunes(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "jobs-stats.db")
|
||||
repo, err := sqlite.Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = repo.Close() })
|
||||
|
||||
h := New(repo, "secret")
|
||||
now := time.Date(2026, 2, 7, 12, 0, 0, 0, time.UTC)
|
||||
nowMs := now.UnixMilli()
|
||||
|
||||
if _, err := repo.DB().Exec(`UPDATE user SET in_flow = 100, out_flow = 200 WHERE id = 1`); err != nil {
|
||||
t.Fatalf("seed user flow: %v", err)
|
||||
}
|
||||
|
||||
if _, err := repo.DB().Exec(`INSERT INTO statistics_flow(user_id, flow, total_flow, time, created_time) VALUES(1, 250, 250, '11:00', ?)`, now.Add(-time.Hour).UnixMilli()); err != nil {
|
||||
t.Fatalf("seed recent statistics row: %v", err)
|
||||
}
|
||||
if _, err := repo.DB().Exec(`INSERT INTO statistics_flow(user_id, flow, total_flow, time, created_time) VALUES(1, 10, 10, '00:00', ?)`, now.Add(-49*time.Hour).UnixMilli()); err != nil {
|
||||
t.Fatalf("seed stale statistics row: %v", err)
|
||||
}
|
||||
|
||||
h.runStatisticsFlowJob(now)
|
||||
|
||||
var staleCount int
|
||||
if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM statistics_flow WHERE created_time < ?`, nowMs-int64((48*time.Hour)/time.Millisecond)).Scan(&staleCount); err != nil {
|
||||
t.Fatalf("query stale statistics rows: %v", err)
|
||||
}
|
||||
if staleCount != 0 {
|
||||
t.Fatalf("expected stale statistics rows to be pruned, got %d", staleCount)
|
||||
}
|
||||
|
||||
var flow int64
|
||||
var total int64
|
||||
var hour string
|
||||
if err := repo.DB().QueryRow(`SELECT flow, total_flow, time FROM statistics_flow WHERE user_id = 1 ORDER BY id DESC LIMIT 1`).Scan(&flow, &total, &hour); err != nil {
|
||||
t.Fatalf("query latest statistics row: %v", err)
|
||||
}
|
||||
if flow != 50 {
|
||||
t.Fatalf("expected increment flow 50, got %d", flow)
|
||||
}
|
||||
if total != 300 {
|
||||
t.Fatalf("expected total flow 300, got %d", total)
|
||||
}
|
||||
if hour != "12:00" {
|
||||
t.Fatalf("expected hour mark 12:00, got %s", hour)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunResetAndExpiryJobResetsFlowAndDisablesExpiredRecords(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "jobs-reset.db")
|
||||
repo, err := sqlite.Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = repo.Close() })
|
||||
|
||||
h := New(repo, "secret")
|
||||
now := time.Date(2026, 3, 15, 0, 0, 5, 0, time.UTC)
|
||||
nowMs := now.UnixMilli()
|
||||
|
||||
if _, err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(2, 'expired_user', 'x', 1, ?, 100, 1000, 2000, 15, 1, ?, ?, 1)
|
||||
`, nowMs-1000, nowMs, nowMs); err != nil {
|
||||
t.Fatalf("insert expired 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, 't1', 1.0, 1, 'tls', 1, ?, ?, 1, NULL, 0)
|
||||
`, nowMs, nowMs); err != nil {
|
||||
t.Fatalf("insert tunnel: %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(10, 2, 1, NULL, 1, 1, 300, 400, 15, ?, 1)
|
||||
`, nowMs-1000); err != nil {
|
||||
t.Fatalf("insert expired user_tunnel: %v", err)
|
||||
}
|
||||
|
||||
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(20, 2, 'expired_user', 'f1', 1, '1.1.1.1:443', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
`, nowMs, nowMs); err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
|
||||
h.runResetAndExpiryJob(now)
|
||||
|
||||
var userIn, userOut int64
|
||||
var userStatus int
|
||||
if err := repo.DB().QueryRow(`SELECT in_flow, out_flow, status FROM user WHERE id = 2`).Scan(&userIn, &userOut, &userStatus); err != nil {
|
||||
t.Fatalf("query user after maintenance: %v", err)
|
||||
}
|
||||
if userIn != 0 || userOut != 0 || userStatus != 0 {
|
||||
t.Fatalf("expected user reset+disabled, got in=%d out=%d status=%d", userIn, userOut, userStatus)
|
||||
}
|
||||
|
||||
var utIn, utOut int64
|
||||
var utStatus int
|
||||
if err := repo.DB().QueryRow(`SELECT in_flow, out_flow, status FROM user_tunnel WHERE id = 10`).Scan(&utIn, &utOut, &utStatus); err != nil {
|
||||
t.Fatalf("query user_tunnel after maintenance: %v", err)
|
||||
}
|
||||
if utIn != 0 || utOut != 0 || utStatus != 0 {
|
||||
t.Fatalf("expected user_tunnel reset+disabled, got in=%d out=%d status=%d", utIn, utOut, utStatus)
|
||||
}
|
||||
|
||||
var forwardStatus int
|
||||
if err := repo.DB().QueryRow(`SELECT status FROM forward WHERE id = 20`).Scan(&forwardStatus); err != nil {
|
||||
t.Fatalf("query forward after maintenance: %v", err)
|
||||
}
|
||||
if forwardStatus != 0 {
|
||||
t.Fatalf("expected forward status=0 after expiry handling, got %d", forwardStatus)
|
||||
}
|
||||
}
|
||||
@@ -56,7 +56,7 @@ 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)
|
||||
roleID := asInt(req["roleId"], asInt(req["role_id"], 1))
|
||||
roleID := 1
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
_, err := db.Exec(`
|
||||
@@ -97,6 +97,20 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
var roleID int
|
||||
if err := db.QueryRow(`SELECT role_id FROM user WHERE id = ?`, id).Scan(&roleID); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
response.WriteJSON(w, response.ErrDefault("用户不存在"))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if roleID == 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("请不要作死"))
|
||||
return
|
||||
}
|
||||
|
||||
var cnt int
|
||||
if err := db.QueryRow(`SELECT COUNT(1) FROM user WHERE user = ? AND id != ?`, username, id).Scan(&cnt); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
@@ -151,6 +165,20 @@ func (h *Handler) userDelete(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
var roleID int
|
||||
if err := h.repo.DB().QueryRow(`SELECT role_id FROM user WHERE id = ?`, id).Scan(&roleID); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
response.WriteJSON(w, response.ErrDefault("用户不存在"))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if roleID == 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("请不要作死"))
|
||||
return
|
||||
}
|
||||
|
||||
db := h.repo.DB()
|
||||
tx, err := db.Begin()
|
||||
if err != nil {
|
||||
@@ -167,6 +195,10 @@ func (h *Handler) userDelete(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if _, err = tx.Exec(`DELETE FROM group_permission_grant WHERE user_tunnel_id IN (SELECT id FROM user_tunnel WHERE user_id = ?)`, id); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if _, err = tx.Exec(`DELETE FROM user_tunnel WHERE user_id = ?`, id); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -175,6 +207,10 @@ func (h *Handler) userDelete(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if _, err = tx.Exec(`DELETE FROM statistics_flow WHERE user_id = ?`, id); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if _, err = tx.Exec(`DELETE FROM user WHERE id = ?`, id); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -248,6 +284,16 @@ func (h *Handler) captchaVerify(w http.ResponseWriter, r *http.Request) {
|
||||
if id == "" {
|
||||
id = asString(req["id"])
|
||||
}
|
||||
trackData := asString(req["data"])
|
||||
if trackData == "" {
|
||||
trackData = asString(req["trackData"])
|
||||
}
|
||||
if id == "" || trackData == "" {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
_, _ = w.Write([]byte(`{"success":false,"message":"bad request"}`))
|
||||
return
|
||||
}
|
||||
h.storeCaptchaToken(id)
|
||||
payload := map[string]interface{}{
|
||||
"success": true,
|
||||
"data": map[string]interface{}{"validToken": id},
|
||||
|
||||
Reference in New Issue
Block a user