mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-05 09:36:37 +08:00
Merge branch 'main' into pr/wss-https-fallback
This commit is contained in:
@@ -18,11 +18,12 @@ func (h *Handler) StartBackgroundJobs() {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
h.jobsCancel = cancel
|
||||
h.jobsStarted = true
|
||||
h.jobsWG.Add(2)
|
||||
h.jobsWG.Add(3)
|
||||
h.jobsMu.Unlock()
|
||||
|
||||
go h.runHourlyStatsLoop(ctx)
|
||||
go h.runDailyMaintenanceLoop(ctx)
|
||||
go h.runNodeRenewalCycleLoop(ctx)
|
||||
}
|
||||
|
||||
func (h *Handler) StopBackgroundJobs() {
|
||||
@@ -176,3 +177,39 @@ func (h *Handler) disableExpiredUserTunnels(nowMs int64) {
|
||||
_ = h.repo.DisableUserTunnel(item.ID)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) runNodeRenewalCycleLoop(ctx context.Context) {
|
||||
defer h.jobsWG.Done()
|
||||
|
||||
for {
|
||||
wait := durationUntilNextNodeRenewalCycle(time.Now())
|
||||
timer := time.NewTimer(wait)
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
if !timer.Stop() {
|
||||
<-timer.C
|
||||
}
|
||||
return
|
||||
case <-timer.C:
|
||||
h.runNodeRenewalCycleJob(time.Now())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func durationUntilNextNodeRenewalCycle(now time.Time) time.Duration {
|
||||
next := now.Truncate(6 * time.Hour).Add(6 * time.Hour)
|
||||
return next.Sub(now)
|
||||
}
|
||||
|
||||
func (h *Handler) runNodeRenewalCycleJob(now time.Time) {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
|
||||
advanced, err := h.repo.AdvanceNodeRenewalCycles(now.UnixMilli())
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
_ = advanced
|
||||
}
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestRunNodeRenewalCycleJob_AdvancesOverdueAnchorTimes(t *testing.T) {
|
||||
dbPath := t.TempDir() + "/renewal-test.db"
|
||||
r, err := repo.Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_ = r.Close()
|
||||
})
|
||||
|
||||
now := time.Date(2026, 3, 8, 12, 0, 0, 0, time.UTC)
|
||||
nowMs := now.UnixMilli()
|
||||
|
||||
nodeID := int64(101)
|
||||
err = r.DB().Exec(`
|
||||
INSERT INTO node (id, name, secret, server_ip, port, http, tls, socks, created_time, status, renewal_cycle, expiry_time)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, nodeID, "no-cycle-node", "test-secret", "192.168.1.1", "1000-65535", 1, 1, 1, nowMs, 1, "", nil).Error
|
||||
if err != nil {
|
||||
t.Fatalf("insert test node: %v", err)
|
||||
}
|
||||
|
||||
quarterNodeID := int64(102)
|
||||
err = r.DB().Exec(`
|
||||
INSERT INTO node (id, name, secret, server_ip, port, http, tls, socks, created_time, status, renewal_cycle, expiry_time)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, quarterNodeID, "quarter-node", "test-secret", "192.168.1.1", "1000-65535", 1, 1, 1, nowMs, 1, "quarter", now.AddDate(0, -4, 0).UnixMilli()).Error
|
||||
if err != nil {
|
||||
t.Fatalf("insert test node: %v", err)
|
||||
}
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.runNodeRenewalCycleJob(now)
|
||||
|
||||
var anchor sql.NullInt64
|
||||
err = r.DB().Raw(`SELECT expiry_time FROM node WHERE id = ?`, quarterNodeID).Row().Scan(&anchor)
|
||||
if err != nil {
|
||||
t.Fatalf("query expiry_time: %v", err)
|
||||
}
|
||||
|
||||
expectedAnchor := now.AddDate(0, 2, 0).UnixMilli()
|
||||
if !anchor.Valid || anchor.Int64 != expectedAnchor {
|
||||
t.Fatalf("expected anchor %d (2026-05-08), got %d", expectedAnchor, anchor.Int64)
|
||||
}
|
||||
}
|
||||
@@ -264,6 +264,7 @@ func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) {
|
||||
nullableText(strings.TrimSpace(asString(req["remark"]))),
|
||||
nullableText(strings.TrimSpace(asString(req["tags"]))),
|
||||
nullableUnixMilli(asInt64(req["expiryTime"], 0)),
|
||||
nullableText(normalizeNodeRenewalCycle(asString(req["renewalCycle"]))),
|
||||
asInt(req["http"], 0),
|
||||
asInt(req["tls"], 0),
|
||||
asInt(req["socks"], 0),
|
||||
@@ -332,6 +333,7 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
nullableText(strings.TrimSpace(asString(req["remark"]))),
|
||||
nullableText(strings.TrimSpace(asString(req["tags"]))),
|
||||
nullableUnixMilli(asInt64(req["expiryTime"], 0)),
|
||||
nullableText(normalizeNodeRenewalCycle(asString(req["renewalCycle"]))),
|
||||
newHTTP,
|
||||
newTLS,
|
||||
newSocks,
|
||||
@@ -3547,6 +3549,15 @@ func nullableUnixMilli(v int64) interface{} {
|
||||
return v
|
||||
}
|
||||
|
||||
func normalizeNodeRenewalCycle(v string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(v)) {
|
||||
case "month", "quarter", "year":
|
||||
return strings.ToLower(strings.TrimSpace(v))
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func nullableInt(v *int64) interface{} {
|
||||
if v == nil {
|
||||
return nil
|
||||
|
||||
@@ -63,6 +63,7 @@ type Node struct {
|
||||
Remark sql.NullString `gorm:"column:remark;type:text"`
|
||||
Tags sql.NullString `gorm:"column:tags;type:text"`
|
||||
ExpiryTime sql.NullInt64 `gorm:"column:expiry_time"`
|
||||
RenewalCycle sql.NullString `gorm:"column:renewal_cycle;type:varchar(20)"`
|
||||
Secret string `gorm:"type:varchar(100);not null"`
|
||||
ServerIP string `gorm:"column:server_ip;type:varchar(100);not null"`
|
||||
ServerIPV4 sql.NullString `gorm:"column:server_ip_v4;type:varchar(100)"`
|
||||
@@ -342,6 +343,7 @@ type NodeBackup struct {
|
||||
Remark string `json:"remark,omitempty"`
|
||||
Tags string `json:"tags,omitempty"`
|
||||
ExpiryTime int64 `json:"expiryTime,omitempty"`
|
||||
RenewalCycle string `json:"renewalCycle,omitempty"`
|
||||
Secret string `json:"secret"`
|
||||
ServerIP string `json:"serverIp"`
|
||||
ServerIPv4 string `json:"serverIpV4,omitempty"`
|
||||
|
||||
@@ -260,7 +260,7 @@ func prepareSQLiteLegacyColumns(db *gorm.DB) error {
|
||||
m := db.Migrator()
|
||||
|
||||
if m.HasTable(&model.Node{}) {
|
||||
for _, field := range []string{"ServerIPV4", "ServerIPV6", "ExtraIPs", "TCPListenAddr", "UDPListenAddr", "Inx", "IsRemote", "RemoteURL", "RemoteToken", "RemoteConfig"} {
|
||||
for _, field := range []string{"ServerIPV4", "ServerIPV6", "ExtraIPs", "TCPListenAddr", "UDPListenAddr", "Inx", "IsRemote", "RemoteURL", "RemoteToken", "RemoteConfig", "Remark", "Tags", "ExpiryTime", "RenewalCycle"} {
|
||||
if m.HasColumn(&model.Node{}, field) {
|
||||
continue
|
||||
}
|
||||
@@ -634,10 +634,11 @@ func (r *Repository) ListNodes() ([]map[string]interface{}, error) {
|
||||
for _, n := range nodes {
|
||||
items = append(items, map[string]interface{}{
|
||||
"id": n.ID, "inx": n.Inx, "name": n.Name,
|
||||
"remark": nullableString(n.Remark),
|
||||
"tags": nullableString(n.Tags),
|
||||
"expiryTime": nullableInt64(n.ExpiryTime),
|
||||
"ip": n.ServerIP, "serverIp": n.ServerIP,
|
||||
"remark": nullableString(n.Remark),
|
||||
"tags": nullableString(n.Tags),
|
||||
"expiryTime": nullableInt64(n.ExpiryTime),
|
||||
"renewalCycle": nullableString(n.RenewalCycle),
|
||||
"ip": n.ServerIP, "serverIp": n.ServerIP,
|
||||
"serverIpV4": nullableString(n.ServerIPV4),
|
||||
"serverIpV6": nullableString(n.ServerIPV6),
|
||||
"extraIPs": nullableString(n.ExtraIPs),
|
||||
@@ -1819,7 +1820,7 @@ func (r *Repository) exportNodes() ([]model.NodeBackup, error) {
|
||||
for _, n := range nodes {
|
||||
b := model.NodeBackup{
|
||||
ID: n.ID, Name: n.Name, Secret: n.Secret, ServerIP: n.ServerIP,
|
||||
Remark: n.Remark.String, Tags: n.Tags.String,
|
||||
Remark: n.Remark.String, Tags: n.Tags.String, RenewalCycle: n.RenewalCycle.String,
|
||||
Port: n.Port, HTTP: n.HTTP, TLS: n.TLS, Socks: n.Socks,
|
||||
CreatedTime: n.CreatedTime, Status: n.Status,
|
||||
TCPListenAddr: n.TCPListenAddr, UDPListenAddr: n.UDPListenAddr,
|
||||
@@ -2180,6 +2181,7 @@ func importNodes(tx *gorm.DB, nodes []model.NodeBackup, now int64) (int, error)
|
||||
Remark: sql.NullString{String: n.Remark, Valid: n.Remark != ""},
|
||||
Tags: sql.NullString{String: n.Tags, Valid: n.Tags != ""},
|
||||
ExpiryTime: sql.NullInt64{Int64: n.ExpiryTime, Valid: n.ExpiryTime > 0},
|
||||
RenewalCycle: sql.NullString{String: n.RenewalCycle, Valid: n.RenewalCycle != ""},
|
||||
Secret: n.Secret,
|
||||
ServerIP: n.ServerIP,
|
||||
ServerIPV4: sql.NullString{String: n.ServerIPv4, Valid: true},
|
||||
@@ -2204,7 +2206,7 @@ func importNodes(tx *gorm.DB, nodes []model.NodeBackup, now int64) (int, error)
|
||||
err := tx.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "id"}},
|
||||
DoUpdates: clause.AssignmentColumns([]string{
|
||||
"name", "remark", "tags", "expiry_time", "secret", "server_ip", "server_ip_v4", "server_ip_v6", "port", "interface_name", "version",
|
||||
"name", "remark", "tags", "expiry_time", "renewal_cycle", "secret", "server_ip", "server_ip_v4", "server_ip_v6", "port", "interface_name", "version",
|
||||
"http", "tls", "socks", "updated_time", "status", "tcp_listen_addr", "udp_listen_addr",
|
||||
"inx", "is_remote", "remote_url", "remote_token", "remote_config",
|
||||
}),
|
||||
|
||||
@@ -7,10 +7,57 @@ import (
|
||||
"testing"
|
||||
|
||||
gsqlite "github.com/glebarez/sqlite"
|
||||
"go-backend/internal/store/model"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/logger"
|
||||
)
|
||||
|
||||
func TestPrepareSQLiteLegacyColumnsAddsNodeMetadataColumns(t *testing.T) {
|
||||
db, err := gorm.Open(gsqlite.Open(":memory:"), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
sqlDB, _ := db.DB()
|
||||
if sqlDB != nil {
|
||||
_ = sqlDB.Close()
|
||||
}
|
||||
})
|
||||
|
||||
if err := db.Exec(`
|
||||
CREATE TABLE node (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
name VARCHAR(100) NOT NULL,
|
||||
secret VARCHAR(100) NOT NULL,
|
||||
server_ip VARCHAR(100) NOT NULL,
|
||||
port TEXT NOT NULL,
|
||||
interface_name VARCHAR(200),
|
||||
version VARCHAR(100),
|
||||
http INTEGER NOT NULL DEFAULT 0,
|
||||
tls INTEGER NOT NULL DEFAULT 0,
|
||||
socks INTEGER NOT NULL DEFAULT 0,
|
||||
created_time INTEGER NOT NULL,
|
||||
updated_time INTEGER,
|
||||
status INTEGER NOT NULL
|
||||
)
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("create legacy node table: %v", err)
|
||||
}
|
||||
|
||||
if err := prepareSQLiteLegacyColumns(db); err != nil {
|
||||
t.Fatalf("prepareSQLiteLegacyColumns: %v", err)
|
||||
}
|
||||
|
||||
m := db.Migrator()
|
||||
for _, field := range []string{"Remark", "Tags", "ExpiryTime", "RenewalCycle"} {
|
||||
if !m.HasColumn(&model.Node{}, field) {
|
||||
t.Fatalf("expected node.%s column to exist", field)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMigrateSchemaRunsPostgresIDRepairEvenAtCurrentVersion(t *testing.T) {
|
||||
db, err := gorm.Open(gsqlite.Open(":memory:"), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
|
||||
@@ -3,6 +3,7 @@ package repo
|
||||
import (
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
@@ -196,7 +197,7 @@ func (r *Repository) GetUserDefaultsForTunnel(userID int64) (flow int64, num int
|
||||
return user.Flow, user.Num, user.ExpTime, user.FlowResetTime, nil
|
||||
}
|
||||
|
||||
func (r *Repository) CreateNode(name, secret, serverIP string, serverIPV4, serverIPV6, port, interfaceName, version, remark, tags, expiryTime interface{}, httpFlag, tlsFlag, socksFlag int, now int64, status int, tcpAddr, udpAddr string, inx, isRemote int, remoteURL, remoteToken, remoteConfig, extraIPs interface{}) error {
|
||||
func (r *Repository) CreateNode(name, secret, serverIP string, serverIPV4, serverIPV6, port, interfaceName, version, remark, tags, expiryTime, renewalCycle interface{}, httpFlag, tlsFlag, socksFlag int, now int64, status int, tcpAddr, udpAddr string, inx, isRemote int, remoteURL, remoteToken, remoteConfig, extraIPs interface{}) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
@@ -205,6 +206,7 @@ func (r *Repository) CreateNode(name, secret, serverIP string, serverIPV4, serve
|
||||
Remark: nullStringFromInterface(remark),
|
||||
Tags: nullStringFromInterface(tags),
|
||||
ExpiryTime: nullInt64FromInterface(expiryTime),
|
||||
RenewalCycle: nullStringFromInterface(renewalCycle),
|
||||
Secret: secret,
|
||||
ServerIP: serverIP,
|
||||
ServerIPV4: nullStringFromInterface(serverIPV4),
|
||||
@@ -242,7 +244,7 @@ func (r *Repository) GetNodeStatusFields(nodeID int64) (status, httpFlag, tlsFla
|
||||
return node.Status, node.HTTP, node.TLS, node.Socks, nil
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateNode(id int64, name, serverIP string, serverIPV4, serverIPV6, port, interfaceName, extraIPs, remark, tags, expiryTime interface{}, httpFlag, tlsFlag, socksFlag int, tcpAddr, udpAddr string, now int64) error {
|
||||
func (r *Repository) UpdateNode(id int64, name, serverIP string, serverIPV4, serverIPV6, port, interfaceName, extraIPs, remark, tags, expiryTime, renewalCycle interface{}, httpFlag, tlsFlag, socksFlag int, tcpAddr, udpAddr string, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
@@ -253,6 +255,7 @@ func (r *Repository) UpdateNode(id int64, name, serverIP string, serverIPV4, ser
|
||||
"remark": nullStringFromInterface(remark),
|
||||
"tags": nullStringFromInterface(tags),
|
||||
"expiry_time": nullInt64FromInterface(expiryTime),
|
||||
"renewal_cycle": nullStringFromInterface(renewalCycle),
|
||||
"server_ip": serverIP,
|
||||
"server_ip_v4": nullStringFromInterface(serverIPV4),
|
||||
"server_ip_v6": nullStringFromInterface(serverIPV6),
|
||||
@@ -1483,3 +1486,59 @@ func (r *Repository) ReplaceUserGroupsByUserID(userID int64, newGroupIDs []int64
|
||||
}
|
||||
return affectedGroupIDs, nil
|
||||
}
|
||||
|
||||
func (r *Repository) AdvanceNodeRenewalCycles(now int64) (int, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
var nodes []model.Node
|
||||
if err := r.db.Where("renewal_cycle IS NOT NULL AND renewal_cycle != '' AND expiry_time IS NOT NULL").Find(&nodes).Error; err != nil {
|
||||
return 0, fmt.Errorf("list nodes with renewal cycle: %w", err)
|
||||
}
|
||||
|
||||
advanced := 0
|
||||
for _, node := range nodes {
|
||||
if !node.ExpiryTime.Valid || node.ExpiryTime.Int64 <= 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
cycleMonths := 0
|
||||
switch node.RenewalCycle.String {
|
||||
case "month":
|
||||
cycleMonths = 1
|
||||
case "quarter":
|
||||
cycleMonths = 3
|
||||
case "year":
|
||||
cycleMonths = 12
|
||||
default:
|
||||
continue
|
||||
}
|
||||
|
||||
anchorTime := node.ExpiryTime.Int64
|
||||
for anchorTime <= now {
|
||||
nextAnchor := advanceByMonths(anchorTime, cycleMonths)
|
||||
if nextAnchor <= anchorTime {
|
||||
break
|
||||
}
|
||||
anchorTime = nextAnchor
|
||||
}
|
||||
|
||||
if anchorTime == node.ExpiryTime.Int64 {
|
||||
continue
|
||||
}
|
||||
|
||||
if err := r.db.Model(&model.Node{}).Where("id = ?", node.ID).Update("expiry_time", anchorTime).Error; err != nil {
|
||||
continue
|
||||
}
|
||||
advanced++
|
||||
}
|
||||
|
||||
return advanced, nil
|
||||
}
|
||||
|
||||
func advanceByMonths(timestamp int64, months int) int64 {
|
||||
t := time.Unix(timestamp/1000, 0)
|
||||
next := t.AddDate(0, months, 0)
|
||||
return next.UnixMilli()
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user