refactor(backend): migrate to modular repository pattern with separated concerns

- Extract database layer into model and repo packages
- Split repository into focused modules (control, federation, flow, groups, mutations)
- Remove monolithic db.go and sqlite/repository.go
- Update handlers to use new repository structure
- Migrate contract tests to new patterns
- Add migration plan documentation
This commit is contained in:
Antigravity
2026-02-17 04:47:11 +00:00
parent 98b4d78b4d
commit 66be07750f
44 changed files with 6829 additions and 6188 deletions
+5 -4
View File
@@ -4,7 +4,7 @@
## OVERVIEW
HTTP request handlers for FLVX Admin API. Core business logic layer.
**Stack:** Go 1.23, net/http, raw SQL (no ORM).
**Stack:** Go 1.23, net/http, GORM via Repository pattern.
## STRUCTURE
```
@@ -28,13 +28,14 @@ handler/
| **Background Jobs** | `jobs.go` | Scheduled sync/cleanup tasks |
## CONVENTIONS
- Inherits from parent: raw SQL, no ORM, JWT in Authorization header.
- Inherits from parent: GORM via Repository pattern, JWT in Authorization header.
- Large files expected (`mutations.go` 3716 LOC - central mutation hub).
- Uses `sqlite.Repository` for DB access via `repo.XXX()` methods.
- Uses `repo.Repository` for DB access via `h.repo.XXX()` methods.
- Handlers never call `repo.DB()` directly — all queries go through Repository methods.
- Domain-driven file split: one file per functional area (federation, jobs, etc.).
## ANTI-PATTERNS
- Do NOT add ORM here - uses raw SQL throughout.
- Do NOT let handlers call `repo.DB()` directly — add a Repository method instead.
- Do NOT change handler signatures without updating router.go.
## COMMANDS
+37 -279
View File
@@ -1,7 +1,6 @@
package handler
import (
"database/sql"
"errors"
"fmt"
"net"
@@ -12,61 +11,18 @@ import (
"time"
"go-backend/internal/http/client"
"go-backend/internal/store/model"
"go-backend/internal/ws"
)
var errForwardNotFound = errors.New("forward not found")
type forwardRecord struct {
ID int64
UserID int64
UserName string
Name string
TunnelID int64
RemoteAddr string
Strategy string
Status int
}
type forwardRecord = model.ForwardRecord
type tunnelRecord = model.TunnelRecord
type forwardPortRecord = model.ForwardPortRecord
type nodeRecord = model.NodeRecord
type tunnelRecord struct {
ID int64
Type int
Status int
Flow int64
TrafficRatio float64
}
type forwardPortRecord struct {
NodeID int64
Port int
}
type nodeRecord struct {
ID int64
Name string
ServerIP string
ServerIPv4 string
ServerIPv6 string
Status int
PortRange string
TCPListenAddr string
UDPListenAddr string
InterfaceName string
IsRemote int
RemoteURL string
RemoteToken string
RemoteConfig string
}
type chainNodeRecord struct {
ChainType int
Inx int64
NodeID int64
Port int
NodeName string
Protocol string
Strategy string
}
type chainNodeRecord = model.ChainNodeRecord
type diagnosisTarget struct {
Address string
@@ -101,247 +57,82 @@ func (h *Handler) ensureTunnelPermission(userID int64, roleID int, tunnelID int6
if roleID == 0 {
return nil
}
var count int
err := h.repo.DB().QueryRow(`SELECT COUNT(1) FROM user_tunnel WHERE user_id = ? AND tunnel_id = ? AND status = 1`, userID, tunnelID).Scan(&count)
ok, err := h.repo.UserTunnelExistsByUserAndTunnel(userID, tunnelID)
if err != nil {
return err
}
if count <= 0 {
if !ok {
return errors.New("你没有该隧道的权限")
}
return nil
}
func (h *Handler) getForwardRecord(forwardID int64) (*forwardRecord, error) {
row := h.repo.DB().QueryRow(`
SELECT id, user_id, user_name, name, tunnel_id, remote_addr, COALESCE(strategy, 'fifo'), status
FROM forward WHERE id = ? LIMIT 1
`, forwardID)
var fr forwardRecord
err := row.Scan(&fr.ID, &fr.UserID, &fr.UserName, &fr.Name, &fr.TunnelID, &fr.RemoteAddr, &fr.Strategy, &fr.Status)
fr, err := h.repo.GetForwardRecord(forwardID)
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, errForwardNotFound
}
return nil, err
}
if strings.TrimSpace(fr.Strategy) == "" {
fr.Strategy = "fifo"
if fr == nil {
return nil, errForwardNotFound
}
return &fr, nil
return fr, nil
}
func (h *Handler) getTunnelRecord(tunnelID int64) (*tunnelRecord, error) {
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, &tr.Flow, &tr.TrafficRatio)
tr, err := h.repo.GetTunnelRecord(tunnelID)
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 == nil {
return nil, errors.New("隧道不存在")
}
if tr.TrafficRatio <= 0 {
tr.TrafficRatio = 1
}
return &tr, nil
return tr, nil
}
func (h *Handler) listForwardsByTunnel(tunnelID int64) ([]forwardRecord, error) {
rows, err := h.repo.DB().Query(`
SELECT id, user_id, user_name, name, tunnel_id, remote_addr, COALESCE(strategy, 'fifo'), status
FROM forward
WHERE tunnel_id = ?
ORDER BY id ASC
`, tunnelID)
if err != nil {
return nil, err
}
defer rows.Close()
result := make([]forwardRecord, 0)
for rows.Next() {
var fr forwardRecord
if err := rows.Scan(&fr.ID, &fr.UserID, &fr.UserName, &fr.Name, &fr.TunnelID, &fr.RemoteAddr, &fr.Strategy, &fr.Status); err != nil {
return nil, err
}
if strings.TrimSpace(fr.Strategy) == "" {
fr.Strategy = "fifo"
}
result = append(result, fr)
}
if err := rows.Err(); err != nil {
return nil, err
}
return result, nil
return h.repo.ListForwardsByTunnel(tunnelID)
}
func (h *Handler) listForwardPorts(forwardID int64) ([]forwardPortRecord, error) {
rows, err := h.repo.DB().Query(`SELECT node_id, port FROM forward_port WHERE forward_id = ? ORDER BY id ASC`, forwardID)
if err != nil {
return nil, err
}
defer rows.Close()
result := make([]forwardPortRecord, 0)
for rows.Next() {
var item forwardPortRecord
if err := rows.Scan(&item.NodeID, &item.Port); err != nil {
return nil, err
}
result = append(result, item)
}
if err := rows.Err(); err != nil {
return nil, err
}
return result, nil
return h.repo.ListForwardPorts(forwardID)
}
func (h *Handler) isTunnelSelectedTLSProtocol(tunnelID int64) (bool, error) {
row := h.repo.DB().QueryRow(`
SELECT protocol
FROM chain_tunnel
WHERE tunnel_id = ? AND chain_type = '3'
ORDER BY id ASC
LIMIT 1
`, tunnelID)
var protocol sql.NullString
if err := row.Scan(&protocol); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return false, nil
}
protocol, err := h.repo.GetTunnelOutProtocol(tunnelID)
if err != nil {
return false, err
}
return isTLSTunnelProtocol(protocol.String), nil
return isTLSTunnelProtocol(protocol), nil
}
func (h *Handler) getNodeRecord(nodeID int64) (*nodeRecord, error) {
row := h.repo.DB().QueryRow(`
SELECT id, name, server_ip, server_ip_v4, server_ip_v6, status, port, tcp_listen_addr, udp_listen_addr, interface_name, is_remote, remote_url, remote_token, remote_config
FROM node
WHERE id = ?
LIMIT 1
`, nodeID)
var n nodeRecord
var serverIPv4 sql.NullString
var serverIPv6 sql.NullString
var portRange sql.NullString
var tcpListen sql.NullString
var udpListen sql.NullString
var iface sql.NullString
var remoteURL sql.NullString
var remoteToken sql.NullString
var remoteConfig sql.NullString
err := row.Scan(&n.ID, &n.Name, &n.ServerIP, &serverIPv4, &serverIPv6, &n.Status, &portRange, &tcpListen, &udpListen, &iface, &n.IsRemote, &remoteURL, &remoteToken, &remoteConfig)
n, err := h.repo.GetNodeRecord(nodeID)
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, errors.New("节点不存在")
}
return nil, err
}
n.ServerIPv4 = strings.TrimSpace(serverIPv4.String)
n.ServerIPv6 = strings.TrimSpace(serverIPv6.String)
n.PortRange = strings.TrimSpace(portRange.String)
n.TCPListenAddr = strings.TrimSpace(tcpListen.String)
n.UDPListenAddr = strings.TrimSpace(udpListen.String)
n.InterfaceName = strings.TrimSpace(iface.String)
n.RemoteURL = strings.TrimSpace(remoteURL.String)
n.RemoteToken = strings.TrimSpace(remoteToken.String)
n.RemoteConfig = strings.TrimSpace(remoteConfig.String)
if n.TCPListenAddr == "" {
n.TCPListenAddr = "[::]"
if n == nil {
return nil, errors.New("节点不存在")
}
if n.UDPListenAddr == "" {
n.UDPListenAddr = "[::]"
}
if strings.TrimSpace(n.Name) == "" {
n.Name = fmt.Sprintf("node_%d", n.ID)
}
return &n, nil
return n, nil
}
func (h *Handler) resolveUserTunnelAndLimiter(userID, tunnelID int64) (int64, *int64, *int, error) {
row := h.repo.DB().QueryRow(`
SELECT ut.id, sl.id, sl.speed
FROM user_tunnel ut
LEFT JOIN speed_limit sl ON sl.id = ut.speed_id
WHERE ut.user_id = ? AND ut.tunnel_id = ?
ORDER BY ut.id ASC
LIMIT 1
`, userID, tunnelID)
var userTunnelID int64
var limiterID sql.NullInt64
var speed sql.NullInt64
err := row.Scan(&userTunnelID, &limiterID, &speed)
info, err := h.repo.ResolveUserTunnelAndLimiter(userID, tunnelID)
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return 0, nil, nil, nil
}
return 0, nil, nil, err
}
if !limiterID.Valid || limiterID.Int64 <= 0 {
return userTunnelID, nil, nil, nil
if info == nil {
return 0, nil, nil, nil
}
v := limiterID.Int64
s := int(speed.Int64)
return userTunnelID, &v, &s, nil
return info.UserTunnelID, info.LimiterID, info.Speed, nil
}
func (h *Handler) listUserTunnelIDs(userID, tunnelID int64) ([]int64, error) {
rows, err := h.repo.DB().Query(`
SELECT id
FROM user_tunnel
WHERE user_id = ? AND tunnel_id = ?
ORDER BY id ASC
`, userID, tunnelID)
if err != nil {
return nil, err
}
defer rows.Close()
out := make([]int64, 0)
for rows.Next() {
var id int64
if err := rows.Scan(&id); err != nil {
return nil, err
}
out = append(out, id)
}
if err := rows.Err(); err != nil {
return nil, err
}
return out, nil
return h.repo.ListUserTunnelIDs(userID, tunnelID)
}
func (h *Handler) listUserTunnelIDsByUser(userID int64) ([]int64, error) {
rows, err := h.repo.DB().Query(`
SELECT id
FROM user_tunnel
WHERE user_id = ?
ORDER BY id ASC
`, userID)
if err != nil {
return nil, err
}
defer rows.Close()
out := make([]int64, 0)
for rows.Next() {
var id int64
if err := rows.Scan(&id); err != nil {
return nil, err
}
out = append(out, id)
}
if err := rows.Err(); err != nil {
return nil, err
}
return out, nil
return h.repo.ListUserTunnelIDsByUser(userID)
}
func (h *Handler) syncForwardServices(forward *forwardRecord, method string, allowFallbackAdd bool) error {
@@ -663,13 +454,13 @@ func (h *Handler) diagnoseTunnelRuntime(tunnelID int64) (map[string]interface{},
return nil, err
}
var tunnelName string
if err := h.repo.DB().QueryRow(`SELECT name FROM tunnel WHERE id = ?`, tunnelID).Scan(&tunnelName); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, errors.New("隧道不存在")
}
tunnelName, err := h.repo.GetTunnelName(tunnelID)
if err != nil {
return nil, err
}
if tunnelName == "" {
return nil, errors.New("隧道不存在")
}
chainRows, err := h.listChainNodesForTunnel(tunnelID)
if err != nil {
@@ -960,40 +751,7 @@ func firstPortFromRange(portRange string) int {
}
func (h *Handler) listChainNodesForTunnel(tunnelID int64) ([]chainNodeRecord, error) {
rows, err := h.repo.DB().Query(`
SELECT CAST(ct.chain_type AS INTEGER), COALESCE(ct.inx, 0), ct.node_id, COALESCE(ct.port, 0), n.name, ct.protocol, ct.strategy
FROM chain_tunnel ct
LEFT JOIN node n ON n.id = ct.node_id
WHERE ct.tunnel_id = ?
ORDER BY CAST(ct.chain_type AS INTEGER) ASC, COALESCE(ct.inx, 0) ASC, ct.id ASC
`, tunnelID)
if err != nil {
return nil, err
}
defer rows.Close()
result := make([]chainNodeRecord, 0)
for rows.Next() {
var item chainNodeRecord
var name sql.NullString
var protocol sql.NullString
var strategy sql.NullString
if err := rows.Scan(&item.ChainType, &item.Inx, &item.NodeID, &item.Port, &name, &protocol, &strategy); err != nil {
return nil, err
}
if strings.TrimSpace(name.String) == "" {
item.NodeName = fmt.Sprintf("node_%d", item.NodeID)
} else {
item.NodeName = name.String
}
item.Protocol = defaultString(protocol.String, "tls")
item.Strategy = defaultString(strategy.String, "round")
result = append(result, item)
}
if err := rows.Err(); err != nil {
return nil, err
}
return result, nil
return h.repo.ListChainNodesForTunnel(tunnelID)
}
func (h *Handler) tcpPingViaNode(nodeID int64, ip string, port int) (map[string]interface{}, error) {
+49 -141
View File
@@ -1,7 +1,6 @@
package handler
import (
"database/sql"
"encoding/json"
"fmt"
"net"
@@ -14,7 +13,7 @@ import (
"go-backend/internal/http/client"
"go-backend/internal/http/response"
"go-backend/internal/store/sqlite"
"go-backend/internal/store/repo"
)
type federationTunnelRequest struct {
@@ -108,7 +107,7 @@ type peerShareUsedPort struct {
}
type peerShareListItem struct {
sqlite.PeerShare
repo.PeerShare
UsedPorts []int `json:"usedPorts"`
UsedPortDetails []peerShareUsedPort `json:"usedPortDetails"`
ActiveRuntimeNum int `json:"activeRuntimeNum"`
@@ -264,7 +263,7 @@ func (h *Handler) federationShareCreate(w http.ResponseWriter, r *http.Request)
now := time.Now().UnixMilli()
token := randomToken(32)
share := &sqlite.PeerShare{
share := &repo.PeerShare{
Name: req.Name,
NodeID: req.NodeID,
Token: token,
@@ -430,40 +429,25 @@ func (h *Handler) federationRemoteUsageList(w http.ResponseWriter, r *http.Reque
return
}
rows, err := h.repo.DB().Query(`
SELECT id, name, remote_url, remote_token, remote_config
FROM node
WHERE is_remote = 1
ORDER BY id DESC
`)
remoteNodes, err := h.repo.ListRemoteNodes()
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
defer rows.Close()
fc := client.NewFederationClient()
localDomain := h.federationLocalDomain()
items := make([]remoteUsageNodeItem, 0)
for rows.Next() {
var (
nodeID int64
nodeName string
remoteURL sql.NullString
remoteToken sql.NullString
remoteConfig sql.NullString
)
if err := rows.Scan(&nodeID, &nodeName, &remoteURL, &remoteToken, &remoteConfig); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
for _, node := range remoteNodes {
nodeID := node.ID
nodeName := node.Name
shareID, maxBandwidth, currentFlow, expiryTime, portRangeStart, portRangeEnd := parseRemoteShareUsageConfig(remoteConfig.String)
shareID, maxBandwidth, currentFlow, expiryTime, portRangeStart, portRangeEnd := parseRemoteShareUsageConfig(node.RemoteConfig.String)
var syncError string
url := strings.TrimSpace(remoteURL.String)
token := strings.TrimSpace(remoteToken.String)
url := strings.TrimSpace(node.RemoteURL.String)
token := strings.TrimSpace(node.RemoteToken.String)
if url != "" && token != "" {
info, connectErr := fc.Connect(url, token, localDomain)
if connectErr != nil {
@@ -484,42 +468,34 @@ func (h *Handler) federationRemoteUsageList(w http.ResponseWriter, r *http.Reque
"portRangeStart": info.PortRangeStart,
"portRangeEnd": info.PortRangeEnd,
})
_, _ = h.repo.DB().Exec(`UPDATE node SET remote_config = ? WHERE id = ?`, string(configData), nodeID)
_ = h.repo.UpdateNodeRemoteConfig(nodeID, string(configData))
}
}
bindingRows, err := h.repo.DB().Query(`
SELECT fb.id, fb.tunnel_id, COALESCE(t.name, ''), fb.chain_type, fb.hop_inx, fb.allocated_port, fb.resource_key, fb.remote_binding_id, fb.updated_time
FROM federation_tunnel_binding fb
LEFT JOIN tunnel t ON t.id = fb.tunnel_id
WHERE fb.node_id = ? AND fb.status = 1
ORDER BY fb.allocated_port ASC, fb.id ASC
`, nodeID)
bindingRows, err := h.repo.ListActiveBindingsForNode(nodeID)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
usedSet := make(map[int]struct{})
bindings := make([]remoteUsageBindingItem, 0)
for bindingRows.Next() {
var item remoteUsageBindingItem
if err := bindingRows.Scan(&item.BindingID, &item.TunnelID, &item.TunnelName, &item.ChainType, &item.HopInx, &item.AllocatedPort, &item.ResourceKey, &item.RemoteBindingID, &item.UpdatedTime); err != nil {
_ = bindingRows.Close()
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
bindings = append(bindings, item)
if item.AllocatedPort > 0 {
usedSet[item.AllocatedPort] = struct{}{}
bindings := make([]remoteUsageBindingItem, 0, len(bindingRows))
for _, b := range bindingRows {
bindings = append(bindings, remoteUsageBindingItem{
BindingID: b.ID,
TunnelID: b.TunnelID,
TunnelName: b.TunnelName,
ChainType: b.ChainType,
HopInx: b.HopInx,
AllocatedPort: b.AllocatedPort,
ResourceKey: b.ResourceKey,
RemoteBindingID: b.RemoteBindingID,
UpdatedTime: b.UpdatedTime,
})
if b.AllocatedPort > 0 {
usedSet[b.AllocatedPort] = struct{}{}
}
}
if err := bindingRows.Err(); err != nil {
_ = bindingRows.Close()
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
_ = bindingRows.Close()
usedPorts := make([]int, 0, len(usedSet))
for port := range usedSet {
@@ -543,10 +519,6 @@ func (h *Handler) federationRemoteUsageList(w http.ResponseWriter, r *http.Reque
SyncError: syncError,
})
}
if err := rows.Err(); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OK(items))
}
@@ -639,31 +611,21 @@ func (h *Handler) nodeImport(w http.ResponseWriter, r *http.Request) {
portRange = fmt.Sprintf("%d-%d", info.PortRangeStart, info.PortRangeEnd)
}
db := h.repo.DB()
inx := nextIndex(db, "node")
inx := h.repo.NextIndex("node")
now := time.Now().UnixMilli()
_, err = db.Exec(`
INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, 0, 0, 0, ?, ?, ?, ?, ?, ?, 1, ?, ?, ?)
`,
if err = h.repo.CreateRemoteNode(
fmt.Sprintf("%s (Remote)", info.NodeName),
randomToken(16), // Dummy secret
randomToken(16),
info.ServerIP,
"", "", // v4/v6 unknown, use server_ip
portRange,
"",
"",
now, now,
now,
info.Status,
"[::]", "[::]",
inx,
req.RemoteURL,
req.Token,
string(configBytes),
)
if err != nil {
); err != nil {
response.WriteJSON(w, response.Err(-2, "Database error: "+err.Error()))
return
}
@@ -755,11 +717,7 @@ func (h *Handler) federationConnect(w http.ResponseWriter, r *http.Request) {
return
}
var nodeName string
var serverIP string
var status int
err = h.repo.DB().QueryRow("SELECT name, server_ip, status FROM node WHERE id = ?", share.NodeID).Scan(&nodeName, &serverIP, &status)
nodeInfo, err := h.repo.GetNodeBasicInfo(share.NodeID)
if err != nil {
response.WriteJSON(w, response.Err(-2, "Node not found"))
return
@@ -769,9 +727,9 @@ func (h *Handler) federationConnect(w http.ResponseWriter, r *http.Request) {
"shareId": share.ID,
"shareName": share.Name,
"nodeId": share.NodeID,
"nodeName": nodeName,
"serverIp": serverIP,
"status": status,
"nodeName": nodeInfo.Name,
"serverIp": nodeInfo.ServerIP,
"status": nodeInfo.Status,
"maxBandwidth": share.MaxBandwidth,
"currentFlow": share.CurrentFlow,
"expiryTime": share.ExpiryTime,
@@ -808,45 +766,20 @@ func (h *Handler) federationTunnelCreate(w http.ResponseWriter, r *http.Request)
return
}
tunnelType := 1
tx, err := h.repo.DB().Begin()
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
defer tx.Rollback()
now := time.Now().UnixMilli()
tunnelID, err := tx.ExecReturningID(`INSERT INTO tunnel (name, type, protocol, flow, created_time, updated_time, status, in_ip) VALUES (?, ?, ?, 0, ?, ?, 1, ?)`,
tunnelID, err := h.repo.CreateFederationTunnel(
fmt.Sprintf("Share-%d-Port-%d", share.ID, req.RemotePort),
tunnelType,
1,
req.Protocol,
now,
now,
"",
)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
_, err = tx.Exec(`INSERT INTO chain_tunnel (tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES (?, '1', ?, ?, 'fifo', 0, ?)`,
tunnelID,
share.NodeID,
req.RemotePort,
req.Protocol,
)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := tx.Commit(); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
h.wsServer.SendCommand(share.NodeID, "reload", nil, time.Second*5)
response.WriteJSON(w, response.OK(map[string]interface{}{
@@ -927,7 +860,7 @@ func (h *Handler) federationRuntimeReservePort(w http.ResponseWriter, r *http.Re
return
}
runtime := &sqlite.PeerShareRuntime{
runtime := &repo.PeerShareRuntime{
ShareID: share.ID,
NodeID: share.NodeID,
ReservationID: randomToken(24),
@@ -981,7 +914,7 @@ func (h *Handler) federationRuntimeApplyRole(w http.ResponseWriter, r *http.Requ
return
}
var runtime *sqlite.PeerShareRuntime
var runtime *repo.PeerShareRuntime
if strings.TrimSpace(req.ReservationID) != "" {
runtime, err = h.repo.GetPeerShareRuntimeByReservationID(share.ID, strings.TrimSpace(req.ReservationID))
} else {
@@ -1152,7 +1085,7 @@ func (h *Handler) federationRuntimeReleaseRole(w http.ResponseWriter, r *http.Re
return
}
var runtime *sqlite.PeerShareRuntime
var runtime *repo.PeerShareRuntime
if strings.TrimSpace(req.BindingID) != "" {
runtime, err = h.repo.GetPeerShareRuntimeByBindingID(share.ID, strings.TrimSpace(req.BindingID))
} else if strings.TrimSpace(req.ReservationID) != "" {
@@ -1299,7 +1232,7 @@ func isFederationServiceCommand(commandType string) bool {
}
}
func validateFederationCommandPorts(share *sqlite.PeerShare, data interface{}) error {
func validateFederationCommandPorts(share *repo.PeerShare, data interface{}) error {
if share == nil || (share.PortRangeStart <= 0 && share.PortRangeEnd <= 0) {
return nil
}
@@ -1339,7 +1272,7 @@ func validateFederationCommandPorts(share *sqlite.PeerShare, data interface{}) e
return nil
}
func (h *Handler) pickPeerSharePort(share *sqlite.PeerShare, requestedPort int) (int, error) {
func (h *Handler) pickPeerSharePort(share *repo.PeerShare, requestedPort int) (int, error) {
if share == nil {
return 0, fmt.Errorf("share not found")
}
@@ -1349,29 +1282,13 @@ func (h *Handler) pickPeerSharePort(share *sqlite.PeerShare, requestedPort int)
used := make(map[int]struct{})
rows, err := h.repo.DB().Query(`SELECT port FROM chain_tunnel WHERE node_id = ? AND port IS NOT NULL AND port > 0`, share.NodeID)
nodePorts, err := h.repo.ListUsedPortsOnNode(share.NodeID)
if err != nil {
return 0, err
}
for rows.Next() {
var p sql.NullInt64
if scanErr := rows.Scan(&p); scanErr == nil && p.Valid && p.Int64 > 0 {
used[int(p.Int64)] = struct{}{}
}
for _, p := range nodePorts {
used[p] = struct{}{}
}
_ = rows.Close()
rows, err = h.repo.DB().Query(`SELECT port FROM forward_port WHERE node_id = ? AND port > 0`, share.NodeID)
if err != nil {
return 0, err
}
for rows.Next() {
var p sql.NullInt64
if scanErr := rows.Scan(&p); scanErr == nil && p.Valid && p.Int64 > 0 {
used[int(p.Int64)] = struct{}{}
}
}
_ = rows.Close()
ports, err := h.repo.ListActivePeerShareRuntimePorts(share.ID, share.NodeID)
if err != nil {
@@ -1412,7 +1329,7 @@ func extractBearerToken(r *http.Request) string {
return ""
}
func isPeerShareFlowExceeded(share *sqlite.PeerShare) bool {
func isPeerShareFlowExceeded(share *repo.PeerShare) bool {
if share == nil {
return false
}
@@ -1655,19 +1572,10 @@ func (h *Handler) cleanupFederationTunnels(shareID int64) {
return
}
namePrefix := fmt.Sprintf("Share-%d-Port-", shareID)
rows, err := h.repo.DB().Query(`SELECT id FROM tunnel WHERE name LIKE ?`, namePrefix+"%")
if err != nil {
tunnelIDs, err := h.repo.ListTunnelIDsByNamePrefix(namePrefix)
if err != nil || len(tunnelIDs) == 0 {
return
}
defer rows.Close()
var tunnelIDs []int64
for rows.Next() {
var id int64
if err := rows.Scan(&id); err == nil {
tunnelIDs = append(tunnelIDs, id)
}
}
for _, tid := range tunnelIDs {
_ = h.deleteTunnelByID(tid)
@@ -10,33 +10,33 @@ import (
"time"
"go-backend/internal/http/response"
"go-backend/internal/store/sqlite"
"go-backend/internal/store/repo"
)
func TestPickPeerSharePortUsesRuntimeReservations(t *testing.T) {
repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db"))
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer repo.Close()
defer r.Close()
h := &Handler{repo: repo}
h := &Handler{repo: r}
now := time.Now().UnixMilli()
if _, err := repo.DB().Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, ?, ?, ?, ?, ?, ?)`, 1, 2, 1, 3000, "round", 1, "tls"); err != nil {
if err := r.DB().Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, ?, ?, ?, ?, ?, ?)`, 1, 2, 1, 3000, "round", 1, "tls").Error; err != nil {
t.Fatalf("insert chain_tunnel: %v", err)
}
if _, err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, 1, 1, 3001); err != nil {
if err := r.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, 1, 1, 3001).Error; err != nil {
t.Fatalf("insert forward_port: %v", err)
}
if _, err := repo.DB().Exec(`
if err := r.DB().Exec(`
INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, 77, 1, "res-1", "rk-1", "b-1", "exit", "", "fed_svc_1", "tls", "round", 3002, "", 1, 1, now, now); err != nil {
`, 77, 1, "res-1", "rk-1", "b-1", "exit", "", "fed_svc_1", "tls", "round", 3002, "", 1, 1, now, now).Error; err != nil {
t.Fatalf("insert peer_share_runtime: %v", err)
}
share := &sqlite.PeerShare{
share := &repo.PeerShare{
ID: 77,
NodeID: 1,
PortRangeStart: 3000,
@@ -57,13 +57,13 @@ func TestPickPeerSharePortUsesRuntimeReservations(t *testing.T) {
}
func TestApplyTunnelRuntimeSkipsRemoteChainAndOutNodes(t *testing.T) {
repo, err := sqlite.Open(filepath.Join(t.TempDir(), "rt-skip.db"))
r, err := repo.Open(filepath.Join(t.TempDir(), "rt-skip.db"))
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer repo.Close()
defer r.Close()
h := &Handler{repo: repo}
h := &Handler{repo: r}
now := time.Now().UnixMilli()
for _, n := range []struct {
id int64
@@ -73,10 +73,10 @@ func TestApplyTunnelRuntimeSkipsRemoteChainAndOutNodes(t *testing.T) {
{12, "remote-chain", "10.99.0.2"},
{13, "remote-out", "10.99.0.3"},
} {
if _, err := repo.DB().Exec(`
if err := r.DB().Exec(`
INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, n.id, n.name, n.name+"-secret", n.ip, n.ip, "", "40000-40010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "http://remote-peer", "remote-token"); err != nil {
`, n.id, n.name, n.name+"-secret", n.ip, n.ip, "", "40000-40010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "http://remote-peer", "remote-token").Error; err != nil {
t.Fatalf("insert node %s: %v", n.name, err)
}
}
@@ -112,25 +112,24 @@ func TestApplyTunnelRuntimeSkipsRemoteChainAndOutNodes(t *testing.T) {
}
func TestPrepareTunnelCreateStateRemoteAutoPortDefersToFederation(t *testing.T) {
repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db"))
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer repo.Close()
defer r.Close()
h := &Handler{repo: repo}
h := &Handler{repo: r}
now := time.Now().UnixMilli()
insertNode := func(name string, status int, portRange string, isRemote int) int64 {
res, execErr := repo.DB().Exec(`
if execErr := r.DB().Exec(`
INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, name, name+"-secret", "10.0.0.1", "10.0.0.1", "", portRange, "", "v1", 1, 1, 1, now, now, status, "[::]", "[::]", 0, isRemote, "http://peer", "peer-token", `{"shareId":1}`)
if execErr != nil {
`, name, name+"-secret", "10.0.0.1", "10.0.0.1", "", portRange, "", "v1", 1, 1, 1, now, now, status, "[::]", "[::]", 0, isRemote, "http://peer", "peer-token", `{"shareId":1}`).Error; execErr != nil {
t.Fatalf("insert node %s: %v", name, execErr)
}
id, idErr := res.LastInsertId()
if idErr != nil {
var id int64
if idErr := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); idErr != nil {
t.Fatalf("node id %s: %v", name, idErr)
}
return id
@@ -139,12 +138,12 @@ func TestPrepareTunnelCreateStateRemoteAutoPortDefersToFederation(t *testing.T)
entryID := insertNode("entry", 1, "31000-31010", 0)
remoteOutID := insertNode("remote-out", 1, "30000", 1)
if _, err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, 1, remoteOutID, 30000); err != nil {
if err := r.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, 1, remoteOutID, 30000).Error; err != nil {
t.Fatalf("insert forward_port: %v", err)
}
tx, err := repo.DB().Begin()
if err != nil {
tx := r.DB().Begin()
if tx.Error != nil {
t.Fatalf("begin tx: %v", err)
}
defer tx.Rollback()
@@ -173,25 +172,24 @@ func TestPrepareTunnelCreateStateRemoteAutoPortDefersToFederation(t *testing.T)
}
func TestPrepareTunnelCreateStateAllowsOfflineRemoteMiddleNode(t *testing.T) {
repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db"))
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer repo.Close()
defer r.Close()
h := &Handler{repo: repo}
h := &Handler{repo: r}
now := time.Now().UnixMilli()
insertNode := func(name string, status int, portRange string, isRemote int) int64 {
res, execErr := repo.DB().Exec(`
if execErr := r.DB().Exec(`
INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, name, name+"-secret", "10.0.0.1", "10.0.0.1", "", portRange, "", "v1", 1, 1, 1, now, now, status, "[::]", "[::]", 0, isRemote, "http://peer", "peer-token", `{"shareId":2}`)
if execErr != nil {
`, name, name+"-secret", "10.0.0.1", "10.0.0.1", "", portRange, "", "v1", 1, 1, 1, now, now, status, "[::]", "[::]", 0, isRemote, "http://peer", "peer-token", `{"shareId":2}`).Error; execErr != nil {
t.Fatalf("insert node %s: %v", name, execErr)
}
id, idErr := res.LastInsertId()
if idErr != nil {
var id int64
if idErr := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); idErr != nil {
t.Fatalf("node id %s: %v", name, idErr)
}
return id
@@ -201,8 +199,8 @@ func TestPrepareTunnelCreateStateAllowsOfflineRemoteMiddleNode(t *testing.T) {
remoteMiddleID := insertNode("middle-remote", 0, "33000-33010", 1)
outID := insertNode("out-local", 1, "34000-34010", 0)
tx, err := repo.DB().Begin()
if err != nil {
tx := r.DB().Begin()
if tx.Error != nil {
t.Fatalf("begin tx: %v", err)
}
defer tx.Rollback()
@@ -238,16 +236,16 @@ func TestPrepareTunnelCreateStateAllowsOfflineRemoteMiddleNode(t *testing.T) {
}
func TestFederationRuntimeReservePortRejectsWhenShareFlowExceeded(t *testing.T) {
repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db"))
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer repo.Close()
defer r.Close()
h := &Handler{repo: repo}
h := &Handler{repo: r}
now := time.Now().UnixMilli()
if err := repo.CreatePeerShare(&sqlite.PeerShare{
if err := r.CreatePeerShare(&repo.PeerShare{
Name: "limited-share",
NodeID: 1,
Token: "limited-token",
@@ -12,28 +12,27 @@ import (
"time"
"go-backend/internal/http/response"
"go-backend/internal/store/sqlite"
"go-backend/internal/store/repo"
)
func TestFederationShareCreateRejectsRemoteNode(t *testing.T) {
repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db"))
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
if err != nil {
t.Fatalf("open sqlite: %v", err)
}
t.Cleanup(func() { _ = repo.Close() })
t.Cleanup(func() { _ = r.Close() })
h := New(repo, "test-jwt-secret")
h := New(r, "test-jwt-secret")
now := time.Now().UnixMilli()
insertRes, err := repo.DB().Exec(`
if err := r.DB().Exec(`
INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, "remote-share-node", "remote-share-secret", "10.10.10.1", "10.10.10.1", "", "20000-20010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "http://peer.example", "peer-token", `{"shareId":1}`)
if err != nil {
`, "remote-share-node", "remote-share-secret", "10.10.10.1", "10.10.10.1", "", "20000-20010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "http://peer.example", "peer-token", `{"shareId":1}`).Error; err != nil {
t.Fatalf("insert remote node: %v", err)
}
remoteNodeID, err := insertRes.LastInsertId()
if err != nil {
var remoteNodeID int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&remoteNodeID); err != nil {
t.Fatalf("get remote node id: %v", err)
}
@@ -71,7 +70,7 @@ func TestFederationShareCreateRejectsRemoteNode(t *testing.T) {
}
var shareCount int
if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM peer_share WHERE node_id = ?`, remoteNodeID).Scan(&shareCount); err != nil {
if err := r.DB().Raw(`SELECT COUNT(1) FROM peer_share WHERE node_id = ?`, remoteNodeID).Row().Scan(&shareCount); err != nil {
t.Fatalf("query peer_share count: %v", err)
}
if shareCount != 0 {
@@ -80,24 +79,23 @@ func TestFederationShareCreateRejectsRemoteNode(t *testing.T) {
}
func TestFederationShareCreateRejectsInvalidAllowedIPs(t *testing.T) {
repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db"))
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
if err != nil {
t.Fatalf("open sqlite: %v", err)
}
t.Cleanup(func() { _ = repo.Close() })
t.Cleanup(func() { _ = r.Close() })
h := New(repo, "test-jwt-secret")
h := New(r, "test-jwt-secret")
now := time.Now().UnixMilli()
insertRes, err := repo.DB().Exec(`
if err := r.DB().Exec(`
INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, "local-share-node", "local-share-secret", "10.20.30.40", "10.20.30.40", "", "21000-21010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 0, "", "", "")
if err != nil {
`, "local-share-node", "local-share-secret", "10.20.30.40", "10.20.30.40", "", "21000-21010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 0, "", "", "").Error; err != nil {
t.Fatalf("insert local node: %v", err)
}
localNodeID, err := insertRes.LastInsertId()
if err != nil {
var localNodeID int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&localNodeID); err != nil {
t.Fatalf("get local node id: %v", err)
}
@@ -136,7 +134,7 @@ func TestFederationShareCreateRejectsInvalidAllowedIPs(t *testing.T) {
}
var shareCount int
if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM peer_share WHERE node_id = ?`, localNodeID).Scan(&shareCount); err != nil {
if err := r.DB().Raw(`SELECT COUNT(1) FROM peer_share WHERE node_id = ?`, localNodeID).Row().Scan(&shareCount); err != nil {
t.Fatalf("query peer_share count: %v", err)
}
if shareCount != 0 {
@@ -145,16 +143,16 @@ func TestFederationShareCreateRejectsInvalidAllowedIPs(t *testing.T) {
}
func TestFederationShareListIncludesRemoteUsedPorts(t *testing.T) {
repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db"))
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
if err != nil {
t.Fatalf("open sqlite: %v", err)
}
t.Cleanup(func() { _ = repo.Close() })
t.Cleanup(func() { _ = r.Close() })
h := New(repo, "test-jwt-secret")
h := New(r, "test-jwt-secret")
now := time.Now().UnixMilli()
if err := repo.CreatePeerShare(&sqlite.PeerShare{
if err := r.CreatePeerShare(&repo.PeerShare{
Name: "provider-share",
NodeID: 9,
Token: "share-list-token",
@@ -169,12 +167,12 @@ func TestFederationShareListIncludesRemoteUsedPorts(t *testing.T) {
t.Fatalf("create peer share: %v", err)
}
share, err := repo.GetPeerShareByToken("share-list-token")
share, err := r.GetPeerShareByToken("share-list-token")
if err != nil || share == nil {
t.Fatalf("load peer share: %v", err)
}
if _, err := repo.DB().Exec(`
if err := r.DB().Exec(`
INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?),
(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?),
@@ -183,7 +181,7 @@ func TestFederationShareListIncludesRemoteUsedPorts(t *testing.T) {
share.ID, share.NodeID, "r-1", "rk-1", "b-1", "middle", "fed_chain_1", "fed_svc_1", "tls", "round", 22001, "", 1, 1, now, now,
share.ID, share.NodeID, "r-2", "rk-2", "b-2", "exit", "", "fed_svc_2", "tls", "round", 22002, "", 1, 1, now, now,
share.ID, share.NodeID, "r-3", "rk-3", "", "", "", "", "tls", "round", 22003, "", 0, 0, now, now,
); err != nil {
).Error; err != nil {
t.Fatalf("insert peer_share_runtime rows: %v", err)
}
@@ -238,16 +236,16 @@ func TestFederationShareListIncludesRemoteUsedPorts(t *testing.T) {
}
func TestFederationShareDeleteCleansUpRuntimes(t *testing.T) {
repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db"))
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
if err != nil {
t.Fatalf("open sqlite: %v", err)
}
t.Cleanup(func() { _ = repo.Close() })
t.Cleanup(func() { _ = r.Close() })
h := New(repo, "test-jwt-secret")
h := New(r, "test-jwt-secret")
now := time.Now().UnixMilli()
if err := repo.CreatePeerShare(&sqlite.PeerShare{
if err := r.CreatePeerShare(&repo.PeerShare{
Name: "delete-cleanup-share",
NodeID: 99,
Token: "delete-cleanup-token",
@@ -261,24 +259,24 @@ func TestFederationShareDeleteCleansUpRuntimes(t *testing.T) {
t.Fatalf("create peer share: %v", err)
}
share, err := repo.GetPeerShareByToken("delete-cleanup-token")
share, err := r.GetPeerShareByToken("delete-cleanup-token")
if err != nil || share == nil {
t.Fatalf("load peer share: %v", err)
}
if _, err := repo.DB().Exec(`
if err := r.DB().Exec(`
INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?),
(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`,
share.ID, 99, "dc-r1", "dc-rk1", "dc-b1", "exit", "", "fed_svc_dc1", "tls", "round", 40001, "", 1, 1, now, now,
share.ID, 99, "dc-r2", "dc-rk2", "dc-b2", "middle", "fed_chain_dc2", "fed_svc_dc2", "tls", "round", 40002, "", 1, 1, now, now,
); err != nil {
).Error; err != nil {
t.Fatalf("insert peer_share_runtime rows: %v", err)
}
var runtimeCount int
if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ? AND status = 1`, share.ID).Scan(&runtimeCount); err != nil {
if err := r.DB().Raw(`SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ? AND status = 1`, share.ID).Row().Scan(&runtimeCount); err != nil {
t.Fatalf("count active runtimes before: %v", err)
}
if runtimeCount != 2 {
@@ -307,7 +305,7 @@ func TestFederationShareDeleteCleansUpRuntimes(t *testing.T) {
}
var shareCount int
if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM peer_share WHERE id = ?`, share.ID).Scan(&shareCount); err != nil {
if err := r.DB().Raw(`SELECT COUNT(1) FROM peer_share WHERE id = ?`, share.ID).Row().Scan(&shareCount); err != nil {
t.Fatalf("count peer_share after: %v", err)
}
if shareCount != 0 {
@@ -315,7 +313,7 @@ func TestFederationShareDeleteCleansUpRuntimes(t *testing.T) {
}
var runtimeCountAfter int
if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ?`, share.ID).Scan(&runtimeCountAfter); err != nil {
if err := r.DB().Raw(`SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ?`, share.ID).Row().Scan(&runtimeCountAfter); err != nil {
t.Fatalf("count peer_share_runtime after: %v", err)
}
if runtimeCountAfter != 0 {
@@ -324,19 +322,19 @@ func TestFederationShareDeleteCleansUpRuntimes(t *testing.T) {
}
func TestFederationRemoteUsageListSyncErrorFallback(t *testing.T) {
repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db"))
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
if err != nil {
t.Fatalf("open sqlite: %v", err)
}
t.Cleanup(func() { _ = repo.Close() })
t.Cleanup(func() { _ = r.Close() })
h := New(repo, "test-jwt-secret")
h := New(r, "test-jwt-secret")
now := time.Now().UnixMilli()
if _, err := repo.DB().Exec(`
if err := r.DB().Exec(`
INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, "sync-error-node", "sync-error-secret", "10.50.60.70", "10.50.60.70", "", "32000-32010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "http://unreachable.invalid:9999", "bad-token", `{"shareId":42,"maxBandwidth":5368709120,"currentFlow":999999,"portRangeStart":32000,"portRangeEnd":32010}`); err != nil {
`, "sync-error-node", "sync-error-secret", "10.50.60.70", "10.50.60.70", "", "32000-32010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "http://unreachable.invalid:9999", "bad-token", `{"shareId":42,"maxBandwidth":5368709120,"currentFlow":999999,"portRangeStart":32000,"portRangeEnd":32010}`).Error; err != nil {
t.Fatalf("insert remote node: %v", err)
}
@@ -380,15 +378,15 @@ func TestFederationRemoteUsageListSyncErrorFallback(t *testing.T) {
}
func TestFederationShareResetFlow(t *testing.T) {
repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db"))
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
if err != nil {
t.Fatalf("open sqlite: %v", err)
}
t.Cleanup(func() { _ = repo.Close() })
t.Cleanup(func() { _ = r.Close() })
h := New(repo, "test-jwt-secret")
h := New(r, "test-jwt-secret")
now := time.Now().UnixMilli()
if err := repo.CreatePeerShare(&sqlite.PeerShare{
if err := r.CreatePeerShare(&repo.PeerShare{
Name: "reset-flow-share",
NodeID: 11,
Token: "reset-flow-token",
@@ -402,7 +400,7 @@ func TestFederationShareResetFlow(t *testing.T) {
}); err != nil {
t.Fatalf("create peer share: %v", err)
}
share, err := repo.GetPeerShareByToken("reset-flow-token")
share, err := r.GetPeerShareByToken("reset-flow-token")
if err != nil || share == nil {
t.Fatalf("load peer share: %v", err)
}
@@ -428,7 +426,7 @@ func TestFederationShareResetFlow(t *testing.T) {
t.Fatalf("expected response code 0, got %d (%s)", payload.Code, payload.Msg)
}
updated, err := repo.GetPeerShare(share.ID)
updated, err := r.GetPeerShare(share.ID)
if err != nil || updated == nil {
t.Fatalf("reload peer share: %v", err)
}
@@ -438,47 +436,50 @@ func TestFederationShareResetFlow(t *testing.T) {
}
func TestFederationRemoteUsageList(t *testing.T) {
repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db"))
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
if err != nil {
t.Fatalf("open sqlite: %v", err)
}
t.Cleanup(func() { _ = repo.Close() })
t.Cleanup(func() { _ = r.Close() })
h := New(repo, "test-jwt-secret")
h := New(r, "test-jwt-secret")
now := time.Now().UnixMilli()
resNode, err := repo.DB().Exec(`
if err := r.DB().Exec(`
INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, "remote-consumer-node", "remote-consumer-secret", "10.30.40.50", "10.30.40.50", "", "31000-31010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "http://peer.example", "peer-token", `{"shareId":88,"maxBandwidth":2147483648,"currentFlow":1073741824,"portRangeStart":31000,"portRangeEnd":31010}`)
if err != nil {
`, "remote-consumer-node", "remote-consumer-secret", "10.30.40.50", "10.30.40.50", "", "31000-31010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "http://peer.example", "peer-token", `{"shareId":88,"maxBandwidth":2147483648,"currentFlow":1073741824,"portRangeStart":31000,"portRangeEnd":31010}`).Error; err != nil {
t.Fatalf("insert remote node: %v", err)
}
nodeID, err := resNode.LastInsertId()
if err != nil {
var nodeID int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&nodeID); err != nil {
t.Fatalf("remote node id: %v", err)
}
resTunnelA, err := repo.DB().Exec(`INSERT INTO tunnel(name, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?)`, "consumer-tunnel-a", 2, "tls", 1, now, now, 1, "", 0)
if err != nil {
if err := r.DB().Exec(`INSERT INTO tunnel(name, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?)`, "consumer-tunnel-a", 2, "tls", 1, now, now, 1, "", 0).Error; err != nil {
t.Fatalf("insert tunnel a: %v", err)
}
tunnelAID, _ := resTunnelA.LastInsertId()
var tunnelAID int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelAID); err != nil {
t.Fatal(err)
}
resTunnelB, err := repo.DB().Exec(`INSERT INTO tunnel(name, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?)`, "consumer-tunnel-b", 2, "tls", 1, now, now, 1, "", 0)
if err != nil {
if err := r.DB().Exec(`INSERT INTO tunnel(name, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?)`, "consumer-tunnel-b", 2, "tls", 1, now, now, 1, "", 0).Error; err != nil {
t.Fatalf("insert tunnel b: %v", err)
}
tunnelBID, _ := resTunnelB.LastInsertId()
var tunnelBID int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelBID); err != nil {
t.Fatal(err)
}
if _, err := repo.DB().Exec(`
if err := r.DB().Exec(`
INSERT INTO federation_tunnel_binding(tunnel_id, node_id, chain_type, hop_inx, remote_url, resource_key, remote_binding_id, allocated_port, status, created_time, updated_time)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?),
(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`,
tunnelAID, nodeID, 2, 1, "http://peer.example", "rk-a", "rb-a", 31001, 1, now, now,
tunnelBID, nodeID, 3, 0, "http://peer.example", "rk-b", "rb-b", 31002, 1, now, now,
); err != nil {
).Error; err != nil {
t.Fatalf("insert federation bindings: %v", err)
}
@@ -532,13 +533,13 @@ func TestFederationRemoteUsageList(t *testing.T) {
}
func TestAuthPeerAllowedIPs(t *testing.T) {
repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db"))
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
if err != nil {
t.Fatalf("open sqlite: %v", err)
}
t.Cleanup(func() { _ = repo.Close() })
t.Cleanup(func() { _ = r.Close() })
h := New(repo, "test-jwt-secret")
h := New(r, "test-jwt-secret")
now := time.Now().UnixMilli()
tests := []struct {
@@ -585,7 +586,7 @@ func TestAuthPeerAllowedIPs(t *testing.T) {
for idx, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
token := fmt.Sprintf("share-token-%d", idx)
if err := repo.CreatePeerShare(&sqlite.PeerShare{
if err := r.CreatePeerShare(&repo.PeerShare{
Name: "share-" + tt.name,
NodeID: 1,
Token: token,
+20 -69
View File
@@ -1,7 +1,6 @@
package handler
import (
"database/sql"
"encoding/json"
"strconv"
"strings"
@@ -207,22 +206,18 @@ func (h *Handler) getUserTunnelPolicy(userTunnelID int64) (*userTunnelPolicy, er
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
}
ut, err := h.repo.GetUserTunnelByID(userTunnelID)
if err != nil {
return nil, err
}
return &policy, nil
if ut == nil {
return nil, nil
}
return &userTunnelPolicy{
ID: ut.ID, UserID: ut.UserID, TunnelID: ut.TunnelID,
Flow: ut.Flow, InFlow: ut.InFlow, OutFlow: ut.OutFlow,
ExpTime: ut.ExpTime, Status: ut.Status,
}, nil
}
func (h *Handler) pauseUserForwards(userID int64, now int64) {
@@ -245,60 +240,20 @@ 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)
_ = h.repo.UpdateForwardStatus(forward.ID, 0, now)
}
}
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, COALESCE(strategy, 'fifo'), 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)
return h.repo.ListActiveForwardsByUser(userID)
}
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, COALESCE(strategy, 'fifo'), 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
return h.repo.ListActiveForwardsByUserTunnel(userID, tunnelID)
}
func (h *Handler) cleanNodeConfigs(nodeID int64, rawConfig string) {
if h == nil || h.repo == nil || h.repo.DB() == nil || nodeID <= 0 {
if h == nil || h.repo == nil || nodeID <= 0 {
return
}
if strings.TrimSpace(rawConfig) == "" {
@@ -383,15 +338,13 @@ func (h *Handler) cleanOrphanedLimiters(nodeID int64, limiters []namedConfigItem
}
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
ok, _ := h.repo.TunnelExists(tunnelID)
return ok
}
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
ok, _ := h.repo.ForwardExists(forwardID)
return ok
}
func (h *Handler) speedLimiterExists(name string) bool {
@@ -402,8 +355,6 @@ func (h *Handler) speedLimiterExists(name string) bool {
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
ok, _ := h.repo.SpeedLimitExists(id)
return ok
}
@@ -5,18 +5,18 @@ import (
"testing"
"time"
"go-backend/internal/store/sqlite"
"go-backend/internal/store/repo"
)
func TestProcessFlowItemTracksPeerShareFlowAndEnforcesLimit(t *testing.T) {
repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db"))
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer repo.Close()
defer r.Close()
now := time.Now().UnixMilli()
if err := repo.CreatePeerShare(&sqlite.PeerShare{
if err := r.CreatePeerShare(&repo.PeerShare{
Name: "flow-share",
NodeID: 1,
Token: "flow-share-token",
@@ -30,22 +30,22 @@ func TestProcessFlowItemTracksPeerShareFlowAndEnforcesLimit(t *testing.T) {
}); err != nil {
t.Fatalf("create peer share: %v", err)
}
share, err := repo.GetPeerShareByToken("flow-share-token")
share, err := r.GetPeerShareByToken("flow-share-token")
if err != nil || share == nil {
t.Fatalf("load peer share: %v", err)
}
if _, err := repo.DB().Exec(`
if err := r.DB().Exec(`
INSERT INTO peer_share_runtime(id, share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, 17, share.ID, share.NodeID, "res-17", "rk-17", "17", "exit", "", "fed_svc_17", "tls", "round", 32001, "", 1, 1, now, now); err != nil {
`, 17, share.ID, share.NodeID, "res-17", "rk-17", "17", "exit", "", "fed_svc_17", "tls", "round", 32001, "", 1, 1, now, now).Error; err != nil {
t.Fatalf("insert peer_share_runtime: %v", err)
}
h := &Handler{repo: repo}
h := &Handler{repo: r}
h.processFlowItem(flowItem{N: "fed_svc_17", U: 1200, D: 900})
updatedShare, err := repo.GetPeerShare(share.ID)
updatedShare, err := r.GetPeerShare(share.ID)
if err != nil || updatedShare == nil {
t.Fatalf("reload share: %v", err)
}
@@ -53,7 +53,7 @@ func TestProcessFlowItemTracksPeerShareFlowAndEnforcesLimit(t *testing.T) {
t.Fatalf("expected current_flow=3100, got %d", updatedShare.CurrentFlow)
}
runtime, err := repo.GetPeerShareRuntimeByID(17)
runtime, err := r.GetPeerShareRuntimeByID(17)
if err != nil || runtime == nil {
t.Fatalf("reload runtime: %v", err)
}
+12 -18
View File
@@ -18,12 +18,12 @@ import (
"go-backend/internal/http/middleware"
"go-backend/internal/http/response"
"go-backend/internal/security"
"go-backend/internal/store/sqlite"
"go-backend/internal/store/repo"
"go-backend/internal/ws"
)
type Handler struct {
repo *sqlite.Repository
repo *repo.Repository
jwtSecret string
wsServer *ws.Server
@@ -69,7 +69,7 @@ type flowItem struct {
D int64 `json:"d"`
}
func New(repo *sqlite.Repository, jwtSecret string) *Handler {
func New(repo *repo.Repository, jwtSecret string) *Handler {
return &Handler{
repo: repo,
jwtSecret: jwtSecret,
@@ -426,7 +426,7 @@ func (h *Handler) openAPISubStore(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
if h == nil || h.repo == nil || h.repo.DB() == nil {
if h == nil || h.repo == nil {
response.WriteJSON(w, response.Err(-2, "database unavailable"))
return
}
@@ -469,27 +469,21 @@ func (h *Handler) openAPISubStore(w http.ResponseWriter, r *http.Request) {
return
}
var userID int64
var inFlow int64
var outFlow int64
var flow int64
var expTime int64
err = h.repo.DB().QueryRow(`SELECT user_id, in_flow, out_flow, flow, exp_time FROM user_tunnel WHERE id = ? LIMIT 1`, tunnelID).
Scan(&userID, &inFlow, &outFlow, &flow, &expTime)
ut, err := h.repo.GetUserTunnelByID(tunnelID)
if err != nil {
if err == sql.ErrNoRows {
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
return
}
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if userID != user.ID {
if ut == nil {
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
return
}
if ut.UserID != user.ID {
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
return
}
headerValue = buildSubscriptionHeader(outFlow, inFlow, flow*giga, expTime/1000)
headerValue = buildSubscriptionHeader(ut.OutFlow, ut.InFlow, ut.Flow*giga, ut.ExpTime/1000)
}
w.Header().Set("subscription-userinfo", headerValue)
@@ -1190,7 +1184,7 @@ func (h *Handler) backupExport(w http.ResponseWriter, r *http.Request) {
type backupImportRequest struct {
Types []string `json:"types"`
sqlite.BackupData
repo.BackupData
}
func (h *Handler) backupImport(w http.ResponseWriter, r *http.Request) {
+15 -107
View File
@@ -2,12 +2,11 @@ package handler
import (
"context"
"database/sql"
"time"
)
func (h *Handler) StartBackgroundJobs() {
if h == nil || h.repo == nil || h.repo.DB() == nil {
if h == nil || h.repo == nil {
return
}
@@ -97,47 +96,28 @@ func durationUntilNextDailyMaintenance(now time.Time) time.Duration {
}
func (h *Handler) runStatisticsFlowJob(now time.Time) {
if h == nil || h.repo == nil || h.repo.DB() == nil {
if h == nil || h.repo == 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)
_ = h.repo.PurgeOldStatisticsFlows(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`)
users, err := h.repo.ListAllUserFlowSnapshots()
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
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)
lastTotal, err := h.repo.GetLastStatisticsFlowTotal(user.UserID)
if err == nil && lastTotal.Valid {
increment = currentTotal - lastTotal.Int64
if increment < 0 {
@@ -145,15 +125,12 @@ func (h *Handler) runStatisticsFlowJob(now time.Time) {
}
}
_, _ = db.Exec(`
INSERT INTO statistics_flow(user_id, flow, total_flow, time, created_time)
VALUES(?, ?, ?, ?, ?)
`, user.userID, increment, currentTotal, hourText, createdTime)
_ = h.repo.CreateStatisticsFlow(user.UserID, increment, currentTotal, hourText, createdTime)
}
}
func (h *Handler) runResetAndExpiryJob(now time.Time) {
if h == nil || h.repo == nil || h.repo.DB() == nil {
if h == nil || h.repo == nil {
return
}
@@ -163,108 +140,39 @@ func (h *Handler) runResetAndExpiryJob(now time.Time) {
}
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)
_ = h.repo.ResetUserMonthlyFlow(currentDay, lastDay)
_ = h.repo.ResetUserTunnelMonthlyFlow(currentDay, lastDay)
}
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)
userIDs, err := h.repo.ListExpiredActiveUserIDs(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)
_ = h.repo.DisableUser(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)
items, err := h.repo.ListExpiredActiveUserTunnels(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)
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)
_ = h.repo.DisableUserTunnel(item.ID)
}
}
+23 -23
View File
@@ -5,36 +5,36 @@ import (
"testing"
"time"
"go-backend/internal/store/sqlite"
"go-backend/internal/store/repo"
)
func TestRunStatisticsFlowJobTracksIncrementAndPrunes(t *testing.T) {
dbPath := filepath.Join(t.TempDir(), "jobs-stats.db")
repo, err := sqlite.Open(dbPath)
r, err := repo.Open(dbPath)
if err != nil {
t.Fatalf("open sqlite: %v", err)
}
t.Cleanup(func() { _ = repo.Close() })
t.Cleanup(func() { _ = r.Close() })
h := New(repo, "secret")
h := New(r, "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 {
if err := r.DB().Exec(`UPDATE user SET in_flow = 100, out_flow = 200 WHERE id = 1`).Error; 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 {
if err := r.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()).Error; 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 {
if err := r.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()).Error; 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 {
if err := r.DB().Raw(`SELECT COUNT(1) FROM statistics_flow WHERE created_time < ?`, nowMs-int64((48*time.Hour)/time.Millisecond)).Row().Scan(&staleCount); err != nil {
t.Fatalf("query stale statistics rows: %v", err)
}
if staleCount != 0 {
@@ -44,7 +44,7 @@ func TestRunStatisticsFlowJobTracksIncrementAndPrunes(t *testing.T) {
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 {
if err := r.DB().Raw(`SELECT flow, total_flow, time FROM statistics_flow WHERE user_id = 1 ORDER BY id DESC LIMIT 1`).Row().Scan(&flow, &total, &hour); err != nil {
t.Fatalf("query latest statistics row: %v", err)
}
if flow != 50 {
@@ -60,41 +60,41 @@ func TestRunStatisticsFlowJobTracksIncrementAndPrunes(t *testing.T) {
func TestRunResetAndExpiryJobResetsFlowAndDisablesExpiredRecords(t *testing.T) {
dbPath := filepath.Join(t.TempDir(), "jobs-reset.db")
repo, err := sqlite.Open(dbPath)
r, err := repo.Open(dbPath)
if err != nil {
t.Fatalf("open sqlite: %v", err)
}
t.Cleanup(func() { _ = repo.Close() })
t.Cleanup(func() { _ = r.Close() })
h := New(repo, "secret")
h := New(r, "secret")
now := time.Date(2026, 3, 15, 0, 0, 5, 0, time.UTC)
nowMs := now.UnixMilli()
if _, err := repo.DB().Exec(`
if err := r.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 {
`, nowMs-1000, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert expired user: %v", err)
}
if _, err := repo.DB().Exec(`
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, 't1', 1.0, 1, 'tls', 1, ?, ?, 1, NULL, 0)
`, nowMs, nowMs); err != nil {
`, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
if _, err := repo.DB().Exec(`
if err := r.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 {
`, nowMs-1000).Error; err != nil {
t.Fatalf("insert expired user_tunnel: %v", err)
}
if _, err := repo.DB().Exec(`
if err := r.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 {
`, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert forward: %v", err)
}
@@ -102,7 +102,7 @@ func TestRunResetAndExpiryJobResetsFlowAndDisablesExpiredRecords(t *testing.T) {
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 {
if err := r.DB().Raw(`SELECT in_flow, out_flow, status FROM user WHERE id = 2`).Row().Scan(&userIn, &userOut, &userStatus); err != nil {
t.Fatalf("query user after maintenance: %v", err)
}
if userIn != 0 || userOut != 0 || userStatus != 0 {
@@ -111,7 +111,7 @@ func TestRunResetAndExpiryJobResetsFlowAndDisablesExpiredRecords(t *testing.T) {
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 {
if err := r.DB().Raw(`SELECT in_flow, out_flow, status FROM user_tunnel WHERE id = 10`).Row().Scan(&utIn, &utOut, &utStatus); err != nil {
t.Fatalf("query user_tunnel after maintenance: %v", err)
}
if utIn != 0 || utOut != 0 || utStatus != 0 {
@@ -119,7 +119,7 @@ func TestRunResetAndExpiryJobResetsFlowAndDisablesExpiredRecords(t *testing.T) {
}
var forwardStatus int
if err := repo.DB().QueryRow(`SELECT status FROM forward WHERE id = 20`).Scan(&forwardStatus); err != nil {
if err := r.DB().Raw(`SELECT status FROM forward WHERE id = 20`).Row().Scan(&forwardStatus); err != nil {
t.Fatalf("query forward after maintenance: %v", err)
}
if forwardStatus != 0 {
File diff suppressed because it is too large Load Diff